// 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/enterprise/connectors/test/fake_content_analysis_delegate.h"

#include <optional>

#include "base/memory/ptr_util.h"
#include "base/task/single_thread_task_runner.h"
#include "base/test/bind.h"
#include "base/time/time.h"
#include "chrome/browser/enterprise/connectors/analysis/page_print_request_handler.h"
#include "chrome/browser/enterprise/connectors/common.h"
#include "chrome/browser/enterprise/connectors/test/fake_clipboard_request_handler.h"
#include "chrome/browser/enterprise/connectors/test/fake_files_request_handler.h"
#include "components/enterprise/common/proto/connectors.pb.h"
#include "components/enterprise/connectors/core/cloud_content_scanning/binary_upload_request.h"
#include "components/enterprise/connectors/core/cloud_content_scanning/clipboard_request_handler.h"
#include "components/enterprise/connectors/core/cloud_content_scanning/deep_scanning_utils.h"

namespace enterprise_connectors::test {

namespace {

base::TimeDelta response_delay = base::Seconds(0);

class FakePagePrintRequestHandler : public PagePrintRequestHandler {
 public:
  static std::unique_ptr<PagePrintRequestHandler> Create(
      base::OnceCallback<void(std::unique_ptr<BinaryUploadRequest>)>
          upload_callback,
      ContentAnalysisInfo* content_analysis_info,
      BinaryUploadService* upload_service,
      Profile* profile,
      GURL url,
      const std::string& printer_name,
      const std::string& page_content_type,
      base::ReadOnlySharedMemoryRegion page_region,
      CompletionCallback callback) {
    auto handler = base::WrapUnique(new FakePagePrintRequestHandler(
        content_analysis_info, upload_service, profile, url, printer_name,
        page_content_type, std::move(page_region), std::move(callback)));
    handler->upload_callback_ = std::move(upload_callback);
    return handler;
  }

 protected:
  using PagePrintRequestHandler::PagePrintRequestHandler;

  void UploadForDeepScanning(
      std::unique_ptr<PagePrintAnalysisRequest> request) override {
    std::move(upload_callback_).Run(std::move(request));
  }

 private:
  base::OnceCallback<void(std::unique_ptr<BinaryUploadRequest>)>
      upload_callback_;
};

}  // namespace

ScanRequestUploadResult FakeContentAnalysisDelegate::result_ =
    ScanRequestUploadResult::kSuccess;
bool FakeContentAnalysisDelegate::dialog_shown_ = false;
bool FakeContentAnalysisDelegate::dialog_canceled_ = false;
int64_t FakeContentAnalysisDelegate::total_analysis_requests_count_ = 0;

FakeContentAnalysisDelegate::FakeContentAnalysisDelegate(
    base::RepeatingClosure delete_closure,
    StatusCallback status_callback,
    std::string dm_token,
    content::WebContents* web_contents,
    Data data,
    CompletionCallback callback,
    DeepScanAccessPoint access_point)
    : ContentAnalysisDelegate(web_contents,
                              std::move(data),
                              std::move(callback),
                              access_point),
      delete_closure_(delete_closure),
      status_callback_(status_callback),
      dm_token_(std::move(dm_token)) {}

FakeContentAnalysisDelegate::~FakeContentAnalysisDelegate() {
  if (!delete_closure_.is_null()) {
    delete_closure_.Run();
  }
}

// static
void FakeContentAnalysisDelegate::SetResponseResult(
    ScanRequestUploadResult result) {
  result_ = result;
}

// static
void FakeContentAnalysisDelegate::
    ResetStaticDialogFlagsAndTotalRequestsCount() {
  dialog_shown_ = false;
  dialog_canceled_ = false;
  total_analysis_requests_count_ = 0;
}

// static
bool FakeContentAnalysisDelegate::WasDialogShown() {
  return dialog_shown_;
}

// static
bool FakeContentAnalysisDelegate::WasDialogCanceled() {
  return dialog_canceled_;
}

int FakeContentAnalysisDelegate::GetTotalAnalysisRequestsCount() {
  return total_analysis_requests_count_;
}

// static
std::unique_ptr<ContentAnalysisDelegate> FakeContentAnalysisDelegate::Create(
    base::RepeatingClosure delete_closure,
    StatusCallback status_callback,
    std::string dm_token,
    content::WebContents* web_contents,
    Data data,
    CompletionCallback callback,
    DeepScanAccessPoint access_point) {
  auto ret = std::make_unique<FakeContentAnalysisDelegate>(
      delete_closure, status_callback, std::move(dm_token), web_contents,
      std::move(data), std::move(callback), access_point);
  FilesRequestHandler::SetFactoryForTesting(base::BindRepeating(
      &FakeFilesRequestHandler::Create,
      base::BindRepeating(
          &FakeContentAnalysisDelegate::FakeUploadFileForDeepScanning,
          base::Unretained(ret.get()))));
  PagePrintRequestHandler::SetFactoryForTesting(base::BindRepeating(
      &FakePagePrintRequestHandler::Create,
      base::BindRepeating(
          &FakeContentAnalysisDelegate::FakeUploadPageForDeepScanning,
          base::Unretained(ret.get()))));
  ClipboardRequestHandler::SetFactoryForTesting(base::BindRepeating(
      &FakeClipboardRequestHandler::Create, base::Unretained(ret.get())));
  return ret;
}

// static
void FakeContentAnalysisDelegate::SetResponseDelay(base::TimeDelta delay) {
  response_delay = delay;
}

// static
ContentAnalysisResponse FakeContentAnalysisDelegate::SuccessfulResponse(
    const std::set<std::string>& tags) {
  ContentAnalysisResponse response;

  auto* result = response.mutable_results()->Add();
  result->set_status(ContentAnalysisResponse::Result::SUCCESS);
  for (const std::string& tag : tags) {
    result->set_tag(tag);
  }

  return response;
}

// static
ContentAnalysisResponse FakeContentAnalysisDelegate::MalwareResponse(
    TriggeredRule::Action action) {
  ContentAnalysisResponse response;

  auto* result = response.mutable_results()->Add();
  result->set_status(ContentAnalysisResponse::Result::SUCCESS);
  result->set_tag("malware");

  auto* rule = result->add_triggered_rules();
  rule->set_action(action);

  return response;
}

// static
ContentAnalysisResponse FakeContentAnalysisDelegate::DlpResponse(
    ContentAnalysisResponse::Result::Status status,
    const std::string& rule_name,
    TriggeredRule::Action action) {
  ContentAnalysisResponse response;

  auto* result = response.mutable_results()->Add();
  result->set_status(status);
  result->set_tag("dlp");

  auto* rule = result->add_triggered_rules();
  rule->set_rule_name(rule_name);
  rule->set_action(action);

  return response;
}

// static
ContentAnalysisResponse FakeContentAnalysisDelegate::MalwareAndDlpResponse(
    TriggeredRule::Action malware_action,
    ContentAnalysisResponse::Result::Status dlp_status,
    const std::string& dlp_rule_name,
    TriggeredRule::Action dlp_action) {
  ContentAnalysisResponse response;

  auto* malware_result = response.add_results();
  malware_result->set_status(ContentAnalysisResponse::Result::SUCCESS);
  malware_result->set_tag("malware");
  auto* malware_rule = malware_result->add_triggered_rules();
  malware_rule->set_action(malware_action);

  auto* dlp_result = response.add_results();
  dlp_result->set_status(dlp_status);
  dlp_result->set_tag("dlp");
  auto* dlp_rule = dlp_result->add_triggered_rules();
  dlp_rule->set_rule_name(dlp_rule_name);
  dlp_rule->set_action(dlp_action);

  return response;
}

ContentAnalysisResponse FakeContentAnalysisDelegate::GetStatus(
    const std::string& contents,
    const base::FilePath& path) {
  return status_callback_.Run(contents, path);
}

void FakeContentAnalysisDelegate::Response(
    std::string contents,
    base::FilePath path,
    std::unique_ptr<BinaryUploadRequest> request,
    std::optional<FakeFilesRequestHandler::FakeFileRequestCallback>
        file_request_callback,
    bool is_image_request) {
  auto response = (status_callback_.is_null() ||
                   result_ != ScanRequestUploadResult::kSuccess)
                      ? ContentAnalysisResponse()
                      : status_callback_.Run(contents, path);
  if (request->IsAuthRequest()) {
    TextRequestCallback(CalculateRequestHandlerResult(
        GetDataForTesting().settings, result_, response));
    return;
  }

  switch (request->analysis_connector()) {
    case AnalysisConnector::BULK_DATA_ENTRY:
    case AnalysisConnector::DATA_COPIED:
      if (is_image_request) {
        ImageRequestCallback(CalculateRequestHandlerResult(
            GetDataForTesting().settings, result_, response));
      } else {
        TextRequestCallback(CalculateRequestHandlerResult(
            GetDataForTesting().settings, result_, response));
      }
      break;
    case AnalysisConnector::FILE_ATTACHED:
    case AnalysisConnector::FILE_DOWNLOADED:
      DCHECK(file_request_callback.has_value());
      std::move(file_request_callback.value()).Run(path, result_, response);
      break;
    case AnalysisConnector::PRINT:
      PageRequestCallback(CalculateRequestHandlerResult(
          GetDataForTesting().settings, result_, response));
      break;
    case AnalysisConnector::FILE_TRANSFER:
    case AnalysisConnector::NETWORK_REQUEST:
    case AnalysisConnector::ANALYSIS_CONNECTOR_UNSPECIFIED:
      NOTREACHED();
  }
}

void FakeContentAnalysisDelegate::FakeUploadFileForDeepScanning(
    ScanRequestUploadResult result,
    const base::FilePath& path,
    std::unique_ptr<BinaryUploadRequest> request,
    FakeFilesRequestHandler::FakeFileRequestCallback callback) {
  DCHECK(!path.empty());
  if (GetDataForTesting()
          .settings.cloud_or_local_settings.is_cloud_analysis()) {
    DCHECK_EQ(dm_token_, request->device_token());
  }

  // Increment total analysis request count.
  total_analysis_requests_count_++;

  // Simulate a response.
  base::SingleThreadTaskRunner::GetCurrentDefault()->PostDelayedTask(
      FROM_HERE,
      base::BindOnce(&FakeContentAnalysisDelegate::Response,
                     weakptr_factory_.GetWeakPtr(), std::string(), path,
                     std::move(request), std::move(callback), false),
      response_delay);
}

void FakeContentAnalysisDelegate::FakeUploadPageForDeepScanning(
    std::unique_ptr<BinaryUploadRequest> request) {
  if (GetDataForTesting()
          .settings.cloud_or_local_settings.is_cloud_analysis()) {
    DCHECK_EQ(dm_token_, request->device_token());
  }

  // Increment total analysis request count.
  total_analysis_requests_count_++;

  // Simulate a response.
  base::SingleThreadTaskRunner::GetCurrentDefault()->PostDelayedTask(
      FROM_HERE,
      base::BindOnce(&FakeContentAnalysisDelegate::Response,
                     weakptr_factory_.GetWeakPtr(), std::string(),
                     base::FilePath(), std::move(request), std::nullopt, false),
      response_delay);
}

void FakeContentAnalysisDelegate::FakeUploadClipboardDataForDeepScanning(
    ClipboardRequestHandler::Type type,
    std::unique_ptr<BinaryUploadRequest> request) {
  if (GetDataForTesting()
          .settings.cloud_or_local_settings.is_cloud_analysis()) {
    DCHECK_EQ(dm_token_, request->device_token());
  }

  // For text/image requests, GetRequestData() is synchronous.
  BinaryUploadRequest::Data data;
  request->GetRequestData(base::BindLambdaForTesting(
      [&data](ScanRequestUploadResult, BinaryUploadRequest::Data data_arg) {
        data = std::move(data_arg);
      }));

  // Increment total analysis request count.
  total_analysis_requests_count_++;

  // Simulate a response.
  base::SingleThreadTaskRunner::GetCurrentDefault()->PostDelayedTask(
      FROM_HERE,
      base::BindOnce(&FakeContentAnalysisDelegate::Response,
                     weakptr_factory_.GetWeakPtr(), data.contents,
                     base::FilePath(), std::move(request), std::nullopt,
                     type == ClipboardRequestHandler::Type::kImage),
      response_delay);
}

bool FakeContentAnalysisDelegate::ShowFinalResultInDialog() {
  dialog_shown_ = true;
  return ContentAnalysisDelegate::ShowFinalResultInDialog();
}

bool FakeContentAnalysisDelegate::CancelDialog() {
  dialog_canceled_ = true;
  return ContentAnalysisDelegate::CancelDialog();
}

BinaryUploadService* FakeContentAnalysisDelegate::GetBinaryUploadService() {
  // This class overrides the upload service, so just return null here.
  return nullptr;
}

}  // namespace enterprise_connectors::test
