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

#include "extensions/browser/service_worker/service_worker_test_utils.h"

#include <utility>

#include "base/containers/map_util.h"
#include "content/public/browser/browser_context.h"
#include "content/public/browser/service_worker_context.h"
#include "content/public/browser/storage_partition.h"
#include "extensions/common/constants.h"
#include "extensions/common/extension.h"
#include "testing/gtest/include/gtest/gtest.h"
#include "third_party/blink/public/mojom/service_worker/service_worker_database.mojom-forward.h"

namespace extensions {
namespace service_worker_test_utils {

content::ServiceWorkerContext* GetServiceWorkerContext(
    content::BrowserContext* browser_context) {
  return browser_context->GetDefaultStoragePartition()
      ->GetServiceWorkerContext();
}

testing::AssertionResult StopServiceWorkerForScope(
    content::ServiceWorkerContext* sw_context,
    const GURL& sw_scope,
    const blink::StorageKey& sw_storage_key) {
  std::optional<int64_t> version_id;
  for (const auto& [id, info] : sw_context->GetRunningServiceWorkerInfos()) {
    if (info.scope == sw_scope && info.key == sw_storage_key) {
      version_id = id;
      break;
    }
  }

  if (!version_id.has_value()) {
    return testing::AssertionFailure()
           << "Could not find running Service Worker for scope "
           << sw_scope.spec();
  }

  extensions::service_worker_test_utils::TestServiceWorkerContextObserver
      observer(sw_context);
  observer.SetRunningId(version_id.value());
  sw_context->StopAllServiceWorkersForStorageKey(sw_storage_key);
  observer.WaitForWorkerStopped();

  return testing::AssertionSuccess();
}

namespace {

std::optional<GURL> GetScopeForExtensionID(
    std::optional<ExtensionId> extension_id) {
  if (!extension_id) {
    return std::nullopt;
  }

  return Extension::GetBaseURLFromExtensionId(*extension_id);
}

}  // namespace

// TestServiceWorkerContextObserver
// ----------------------------------------------------

TestServiceWorkerContextObserver::TestServiceWorkerContextObserver(
    content::ServiceWorkerContext* context,
    std::optional<ExtensionId> extension_id)
    : extension_scope_(GetScopeForExtensionID(std::move(extension_id))),
      context_(context) {
  scoped_observation_.Observe(context_);
  scoped_sync_observation_.Observe(context_);
}

TestServiceWorkerContextObserver::TestServiceWorkerContextObserver(
    content::BrowserContext* browser_context,
    std::optional<ExtensionId> extension_id)
    : extension_scope_(GetScopeForExtensionID(std::move(extension_id))),
      context_(GetServiceWorkerContext(browser_context)) {
  scoped_observation_.Observe(context_);
  scoped_sync_observation_.Observe(context_);
}

TestServiceWorkerContextObserver::~TestServiceWorkerContextObserver() = default;

void TestServiceWorkerContextObserver::WaitForRegistrationStored() {
  if (registration_stored_) {
    return;
  }

  SCOPED_TRACE("Waiting for worker registration to be stored");
  base::RunLoop run_loop;
  stored_quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();
}

int64_t TestServiceWorkerContextObserver::WaitForStartWorkerMessageSent() {
  if (!start_message_sent_version_id_) {
    SCOPED_TRACE("Waiting for StartWorker message to be sent");
    base::RunLoop run_loop;
    start_message_sent_quit_closure_ = run_loop.QuitClosure();
    run_loop.Run();
  }

  return *start_message_sent_version_id_;
}

int64_t TestServiceWorkerContextObserver::WaitForWorkerStarted() {
  if (!running_version_id_) {
    SCOPED_TRACE("Waiting for worker to be started");
    base::RunLoop run_loop;
    started_quit_closure_ = run_loop.QuitClosure();
    run_loop.Run();
  }

  return *running_version_id_;
}

int64_t TestServiceWorkerContextObserver::WaitForWorkerStopping() {
  if (!running_version_id_) {
    return blink::mojom::kInvalidServiceWorkerVersionId;
  }
  if (stopping_version_id_) {
    return *stopping_version_id_;
  }

  SCOPED_TRACE("Waiting for worker to be stopping");
  base::RunLoop run_loop;
  stopping_quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();

  return *stopping_version_id_;
}

int64_t TestServiceWorkerContextObserver::WaitForWorkerStopped() {
  if (!running_version_id_) {
    return blink::mojom::kInvalidServiceWorkerVersionId;
  }
  if (stopped_version_id_) {
    return *stopped_version_id_;
  }

  SCOPED_TRACE("Waiting for worker to be stopped");
  base::RunLoop run_loop;
  stopped_quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();

  return *stopped_version_id_;
}

int64_t TestServiceWorkerContextObserver::WaitForWorkerActivated() {
  if (!activated_version_id_) {
    SCOPED_TRACE("Waiting for worker to be activated");
    base::RunLoop run_loop;
    activated_quit_closure_ = run_loop.QuitClosure();
    run_loop.Run();
  }

  return *activated_version_id_;
}

int TestServiceWorkerContextObserver::GetCompletedCount(
    const GURL& scope) const {
  const auto it = registrations_completed_map_.find(scope);
  return it == registrations_completed_map_.end() ? 0 : it->second;
}

void TestServiceWorkerContextObserver::OnRegistrationCompleted(
    const GURL& scope) {
  ++registrations_completed_map_[scope];
}

void TestServiceWorkerContextObserver::OnRegistrationStored(
    int64_t registration_id,
    const GURL& scope,
    const content::ServiceWorkerRegistrationInformation& service_worker_info) {
  if (scope.SchemeIs(kExtensionScheme)) {
    registration_stored_ = true;
    if (stored_quit_closure_) {
      std::move(stored_quit_closure_).Run();
    }
  }
}

void TestServiceWorkerContextObserver::OnStartWorkerMessageSentSync(
    int64_t version_id,
    const GURL& scope) {
  if (extension_scope_ && extension_scope_ != scope) {
    return;
  }

  start_message_sent_version_id_ = version_id;
  if (start_message_sent_quit_closure_) {
    std::move(start_message_sent_quit_closure_).Run();
  }
}

void TestServiceWorkerContextObserver::OnVersionStartedRunning(
    int64_t version_id,
    const content::ServiceWorkerRunningInfo& running_info) {
  if (extension_scope_ && extension_scope_ != running_info.scope) {
    return;
  }

  running_version_id_ = version_id;
  if (started_quit_closure_) {
    std::move(started_quit_closure_).Run();
  }
}

void TestServiceWorkerContextObserver::OnStoppingSync(
    int64_t version_id,
    const GURL& scope,
    const blink::ServiceWorkerToken& service_worker_token) {
  if (running_version_id_ && running_version_id_ == version_id) {
    stopping_version_id_ = version_id;
    if (stopping_quit_closure_) {
      std::move(stopping_quit_closure_).Run();
    }
  }
}

void TestServiceWorkerContextObserver::OnVersionStoppedRunning(
    int64_t version_id) {
  if (running_version_id_ && running_version_id_ == version_id) {
    stopped_version_id_ = version_id;
    if (stopped_quit_closure_) {
      std::move(stopped_quit_closure_).Run();
    }
  }
}

void TestServiceWorkerContextObserver::OnVersionActivated(int64_t version_id,
                                                          const GURL& scope) {
  if (extension_scope_ && extension_scope_ != scope) {
    return;
  }

  activated_version_id_ = version_id;
  if (activated_quit_closure_) {
    std::move(activated_quit_closure_).Run();
  }
}

void TestServiceWorkerContextObserver::OnDestruct(
    content::ServiceWorkerContext* context) {
  scoped_observation_.Reset();
  scoped_sync_observation_.Reset();
  context_ = nullptr;
}

// UnregisterWorkerObserver ----------------------------------------------------
UnregisterWorkerObserver::UnregisterWorkerObserver(
    ProcessManager* process_manager,
    const ExtensionId& extension_id)
    : extension_id_(extension_id) {
  observation_.Observe(process_manager);
}

UnregisterWorkerObserver::~UnregisterWorkerObserver() = default;

void UnregisterWorkerObserver::OnStoppedTrackingServiceWorkerInstance(
    content::BrowserContext& browser_context,
    const WorkerId& worker_id) {
  run_loop_.QuitWhenIdle();
}

void UnregisterWorkerObserver::WaitForUnregister() {
  run_loop_.Run();
}

TestServiceWorkerTaskQueueObserver::TestServiceWorkerTaskQueueObserver() {
  ServiceWorkerTaskQueue::SetObserverForTest(this);
}

TestServiceWorkerTaskQueueObserver::~TestServiceWorkerTaskQueueObserver() {
  ServiceWorkerTaskQueue::SetObserverForTest(nullptr);
}

void TestServiceWorkerTaskQueueObserver::WaitForWorkerStarted(
    const ExtensionId& extension_id) {
  if (started_set_.count(extension_id) != 0) {
    return;
  }

  base::RunLoop run_loop;
  started_quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();
}

void TestServiceWorkerTaskQueueObserver::WaitForWorkerStopped(
    const ExtensionId& extension_id) {
  if (stopped_set_.count(extension_id) != 0) {
    return;
  }

  base::RunLoop run_loop;
  quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();
}

void TestServiceWorkerTaskQueueObserver::WaitForUntrackServiceWorkerState(
    const GURL& scope) {
  if (untracked_set_.count(scope) != 0) {
    return;
  }

  base::RunLoop run_loop;
  untrack_quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();
}

void TestServiceWorkerTaskQueueObserver::WaitForWorkerContextInitialized(
    const ExtensionId& extension_id) {
  if (inited_set_.count(extension_id) != 0) {
    return;
  }

  base::RunLoop run_loop;
  quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();
}

TestServiceWorkerTaskQueueObserver::WorkerStartFailedData
TestServiceWorkerTaskQueueObserver::WaitForDidStartWorkerFail(
    const ExtensionId& extension_id) {
  const WorkerStartFailedData* const data =
      base::FindOrNull(failed_map_, extension_id);
  if (data) {
    return *data;
  }

  base::RunLoop run_loop;
  quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();

  return failed_map_[extension_id];
}

void TestServiceWorkerTaskQueueObserver::WaitForOnActivateExtension(
    const ExtensionId& extension_id) {
  if (activated_map_.count(extension_id) == 1) {
    return;
  }

  base::RunLoop run_loop;
  quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();
}

bool TestServiceWorkerTaskQueueObserver::WaitForRegistrationMismatchMitigation(
    const ExtensionId& extension_id) {
  const bool* const value = base::FindOrNull(mitigated_map_, extension_id);
  if (value) {
    return *value;
  }

  base::RunLoop run_loop;
  quit_closure_ = run_loop.QuitClosure();
  run_loop.Run();

  return mitigated_map_[extension_id];
}

std::optional<bool>
TestServiceWorkerTaskQueueObserver::WillRegisterServiceWorker(
    const ExtensionId& extension_id) const {
  const bool* const value = base::FindOrNull(activated_map_, extension_id);
  if (value) {
    return *value;
  }
  return std::nullopt;
}

int TestServiceWorkerTaskQueueObserver::GetRequestedWorkerStartedCount(
    const ExtensionId& extension_id) const {
  const int* const value =
      base::FindOrNull(requested_worker_started_map_, extension_id);
  return value ? *value : 0;
}

void TestServiceWorkerTaskQueueObserver::DidStartWorker(
    const ExtensionId& extension_id) {
  started_set_.insert(extension_id);
  if (started_quit_closure_) {
    std::move(started_quit_closure_).Run();
  }
}

void TestServiceWorkerTaskQueueObserver::DidStartWorkerFail(
    const ExtensionId& extension_id,
    size_t num_pending_tasks,
    blink::ServiceWorkerStatusCode status_code) {
  WorkerStartFailedData& data = failed_map_[extension_id];
  data.num_pending_tasks = num_pending_tasks;
  data.status_code = status_code;
  if (quit_closure_) {
    std::move(quit_closure_).Run();
  }
}

void TestServiceWorkerTaskQueueObserver::
    RendererDidInitializeServiceWorkerContext(const ExtensionId& extension_id) {
  inited_set_.insert(extension_id);
  if (quit_closure_) {
    std::move(quit_closure_).Run();
  }
}

void TestServiceWorkerTaskQueueObserver::OnActivateExtension(
    const ExtensionId& extension_id,
    bool will_register_service_worker) {
  activated_map_[extension_id] = will_register_service_worker;
  if (quit_closure_) {
    std::move(quit_closure_).Run();
  }
}

void TestServiceWorkerTaskQueueObserver::RegistrationMismatchMitigated(
    const ExtensionId& extension_id,
    bool success) {
  mitigated_map_[extension_id] = success;
  if (quit_closure_) {
    std::move(quit_closure_).Run();
  }
}

void TestServiceWorkerTaskQueueObserver::RequestedWorkerStart(
    const ExtensionId& extension_id) {
  ++requested_worker_started_map_[extension_id];
}

void TestServiceWorkerTaskQueueObserver::RendererDidStopServiceWorkerContext(
    const ExtensionId& extension_id) {
  stopped_set_.insert(extension_id);
  if (quit_closure_) {
    std::move(quit_closure_).Run();
  }
}

void TestServiceWorkerTaskQueueObserver::UntrackServiceWorkerState(
    const GURL& scope) {
  untracked_set_.insert(scope);
  if (untrack_quit_closure_) {
    std::move(untrack_quit_closure_).Run();
  }
}

}  // namespace service_worker_test_utils
}  // namespace extensions
