From 85a01e9636233453828049aa02762d69ac02f282 Mon Sep 17 00:00:00 2001 From: Guogang Li Date: Wed, 12 Jul 2023 13:30:48 -0700 Subject: [PATCH] internal refactor PiperOrigin-RevId: 547586065 --- internal/test/fake_task_runner.cc | 137 +++---------------------- internal/test/fake_task_runner.h | 55 +++------- internal/test/fake_task_runner_test.cc | 48 --------- 3 files changed, 31 insertions(+), 209 deletions(-) diff --git a/internal/test/fake_task_runner.cc b/internal/test/fake_task_runner.cc index cd35a9a1..956832ac 100644 --- a/internal/test/fake_task_runner.cc +++ b/internal/test/fake_task_runner.cc @@ -14,8 +14,6 @@ #include "internal/test/fake_task_runner.h" -#include -#include // NOLINT #include #include #include @@ -23,30 +21,22 @@ #include "absl/synchronization/mutex.h" #include "absl/synchronization/notification.h" #include "absl/time/time.h" +#include "internal/platform/count_down_latch.h" #include "internal/test/fake_timer.h" namespace nearby { std::atomic_uint FakeTaskRunner::total_running_thread_count_ = 0; -FakeTaskRunner::~FakeTaskRunner() { - absl::MutexLock lock(&mutex_); - CleanThreads(); -} +FakeTaskRunner::~FakeTaskRunner() { absl::MutexLock lock(&mutex_); } bool FakeTaskRunner::PostTask(absl::AnyInvocable task) { absl::MutexLock lock(&mutex_); - if (mode_ == Mode::kActive) { - if (running_thread_count_ >= count_) { - queued_tasks_.push_back(std::move(task)); - return true; - } - - ++running_thread_count_; - Run(std::move(task)); - return true; - } - pending_tasks_.push_back(std::move(task)); + ++total_running_thread_count_; + task_executor_->Execute([task = std::move(task)]() mutable { + task(); + --total_running_thread_count_; + }); return true; } @@ -55,67 +45,27 @@ bool FakeTaskRunner::PostDelayedTask(absl::Duration delay, absl::MutexLock lock(&mutex_); std::unique_ptr timer = std::make_unique(clock_); Timer* timer_ptr = timer.get(); - uint32_t id = GenerateId(); - queued_delayed_tasks_.emplace(id, std::move(timer)); - timer_ptr->Start(delay / absl::Milliseconds(1), 0, - [this, task = std::move(task), id]() mutable { - PostTask(std::move(task)); - { - absl::MutexLock lock(&mutex_); - completed_delayed_tasks_.push_back(id); - } - }); + timers_.push_back(std::move(timer)); + timer_ptr->Start( + delay / absl::Milliseconds(1), 0, + [this, task = std::move(task)]() mutable { PostTask(std::move(task)); }); return true; } -void FakeTaskRunner::SetMode(Mode mode) { - absl::MutexLock lock(&mutex_); - mode_ = mode; -} - -FakeTaskRunner::Mode FakeTaskRunner::GetMode() const { - absl::MutexLock lock(&mutex_); - return mode_; -} - -void FakeTaskRunner::RunNextPendingTask() { - absl::MutexLock lock(&mutex_); - InternalRunNextPendingTask(); -} - -void FakeTaskRunner::RunAllPendingTasks() { - absl::MutexLock lock(&mutex_); - while (!pending_tasks_.empty()) { - InternalRunNextPendingTask(); - } -} - void FakeTaskRunner::Sync() { absl::Notification notification; PostTask([&] { notification.Notify(); }); notification.WaitForNotification(); } -const std::vector>& -FakeTaskRunner::GetAllPendingTasks() const { - absl::MutexLock lock(&mutex_); - return pending_tasks_; -} - -const absl::flat_hash_map>& -FakeTaskRunner::GetAllDelayedTasks() { - absl::MutexLock lock(&mutex_); - if (!completed_delayed_tasks_.empty()) { - for (uint32_t id : completed_delayed_tasks_) { - queued_delayed_tasks_.erase(id); - } +bool FakeTaskRunner::SyncWithTimeout(absl::Duration timeout) { + CountDownLatch latch(count_); + for (int i = 0; i < count_; ++i) { + PostTask([&] { latch.CountDown(); }); } - return queued_delayed_tasks_; -} -int FakeTaskRunner::GetConcurrentCount() const { - absl::MutexLock lock(&mutex_); - return count_; + auto result = latch.Await(timeout); + return result.ok() && result.result(); } bool FakeTaskRunner::WaitForRunningTasksWithTimeout(absl::Duration timeout) { @@ -128,57 +78,4 @@ bool FakeTaskRunner::WaitForRunningTasksWithTimeout(absl::Duration timeout) { return total_running_thread_count_ == 0; } -uint32_t FakeTaskRunner::GenerateId() { - ++current_id_; - return current_id_; -} - -void FakeTaskRunner::CleanThreads() { - auto it = threads_.begin(); - while (it != threads_.end()) { - // Delete the thread if it is ready - auto status = it->wait_for(std::chrono::seconds(0)); - if (status == std::future_status::ready) { - it = threads_.erase(it); - } else { - ++it; - } - } -} - -void FakeTaskRunner::Run(absl::AnyInvocable task) { - CleanThreads(); - ++total_running_thread_count_; - // Run the task in a new thread, to simulate the real environment. - std::future thread = - std::async(std::launch::async, [&, task = std::move(task)]() mutable { - task(); - RunNextQueueTask(); - --total_running_thread_count_; - }); - threads_.push_back(std::move(thread)); -} - -void FakeTaskRunner::InternalRunNextPendingTask() { - if (pending_tasks_.empty()) { - return; - } - - Run(std::move(pending_tasks_.front())); - pending_tasks_.erase(pending_tasks_.begin()); -} - -void FakeTaskRunner::RunNextQueueTask() { - absl::MutexLock lock(&mutex_); - --running_thread_count_; - if (queued_tasks_.empty()) { - return; - } - - auto task = std::move(queued_tasks_.front()); - queued_tasks_.erase(queued_tasks_.begin()); - ++running_thread_count_; - Run(std::move(task)); -} - } // namespace nearby diff --git a/internal/test/fake_task_runner.h b/internal/test/fake_task_runner.h index f92b07f9..643e394e 100644 --- a/internal/test/fake_task_runner.h +++ b/internal/test/fake_task_runner.h @@ -17,25 +17,24 @@ #include #include -#include //NOLINT #include #include #include "absl/base/thread_annotations.h" -#include "absl/container/flat_hash_map.h" #include "absl/time/time.h" +#include "internal/platform/multi_thread_executor.h" #include "internal/platform/task_runner.h" +#include "internal/platform/timer.h" #include "internal/test/fake_clock.h" -#include "internal/test/fake_timer.h" namespace nearby { class FakeTaskRunner : public TaskRunner { public: - enum class Mode { kActive, kPending }; - FakeTaskRunner(FakeClock* clock, uint32_t count) - : clock_(clock), count_(count) {} + : clock_(clock), + count_(count), + task_executor_(std::make_unique(count)) {} ~FakeTaskRunner() override ABSL_LOCKS_EXCLUDED(mutex_); bool PostTask(absl::AnyInvocable task) override @@ -47,51 +46,25 @@ class FakeTaskRunner : public TaskRunner { absl::AnyInvocable task) override ABSL_LOCKS_EXCLUDED(mutex_); - // Mocked methods. - void SetMode(Mode mode) ABSL_LOCKS_EXCLUDED(mutex_); - Mode GetMode() const ABSL_LOCKS_EXCLUDED(mutex_); - - void RunNextPendingTask() ABSL_LOCKS_EXCLUDED(mutex_); - void RunAllPendingTasks() ABSL_LOCKS_EXCLUDED(mutex_); + // Wait for all thread completed. void Sync(); - const std::vector>& GetAllPendingTasks() const - ABSL_LOCKS_EXCLUDED(mutex_); - const absl::flat_hash_map>& - GetAllDelayedTasks() ABSL_LOCKS_EXCLUDED(mutex_); + // In some test cases, we only need to wait for a timeout . + bool SyncWithTimeout(absl::Duration timeout); - int GetConcurrentCount() const ABSL_LOCKS_EXCLUDED(mutex_); - - // In some test cases, we needs to make sure all running tasks completion + // In some test cases, we need to make sure all running tasks completion // before go to next task. This method can be used for the purpose. static bool WaitForRunningTasksWithTimeout(absl::Duration timeout); - static int GetTotalRunningThreadCount() { - return total_running_thread_count_; - } private: - uint32_t GenerateId(); - void CleanThreads() ABSL_SHARED_LOCKS_REQUIRED(mutex_); - void Run(absl::AnyInvocable task) ABSL_SHARED_LOCKS_REQUIRED(mutex_); - void InternalRunNextPendingTask() ABSL_SHARED_LOCKS_REQUIRED(mutex_); - void RunNextQueueTask() ABSL_LOCKS_EXCLUDED(mutex_); - mutable absl::Mutex mutex_; - mutable absl::Mutex thread_mutex_; - Mode mode_ ABSL_GUARDED_BY(mutex_) = Mode::kActive; - std::atomic_uint current_id_ = 0; FakeClock* clock_ = nullptr; - uint32_t count_ ABSL_GUARDED_BY(mutex_); + uint32_t count_ = 0; + std::unique_ptr task_executor_ ABSL_GUARDED_BY(mutex_) = + nullptr; - // Used for pending mode - std::vector> pending_tasks_ - ABSL_GUARDED_BY(mutex_); - std::vector> queued_tasks_ ABSL_GUARDED_BY(mutex_); - std::vector completed_delayed_tasks_ ABSL_GUARDED_BY(mutex_); - absl::flat_hash_map> queued_delayed_tasks_ - ABSL_GUARDED_BY(mutex_); - std::vector> threads_ ABSL_GUARDED_BY(mutex_); - int running_thread_count_ ABSL_GUARDED_BY(mutex_) = 0; + // Tracks delayed tasks. + std::vector> timers_ ABSL_GUARDED_BY(mutex_); static std::atomic_uint total_running_thread_count_; }; diff --git a/internal/test/fake_task_runner_test.cc b/internal/test/fake_task_runner_test.cc index e778d0ae..a09b96ae 100644 --- a/internal/test/fake_task_runner_test.cc +++ b/internal/test/fake_task_runner_test.cc @@ -32,7 +32,6 @@ TEST(FakeTaskRunner, PostTask) { task_runner.PostTask([&count] { ++count; }); ASSERT_TRUE( FakeTaskRunner::WaitForRunningTasksWithTimeout(absl::Milliseconds(100))); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 0); EXPECT_EQ(count, 1); } @@ -41,58 +40,11 @@ TEST(FakeTaskRunner, PostDelayedTask) { int count = 0; FakeTaskRunner task_runner{&clock, 1}; task_runner.PostDelayedTask(absl::Seconds(10), [&count] { ++count; }); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 0); - EXPECT_EQ(task_runner.GetAllDelayedTasks().size(), 1); EXPECT_EQ(count, 0); clock.FastForward(absl::Seconds(10)); ASSERT_TRUE( FakeTaskRunner::WaitForRunningTasksWithTimeout(absl::Milliseconds(100))); EXPECT_EQ(count, 1); - EXPECT_EQ(task_runner.GetAllDelayedTasks().size(), 0); -} - -TEST(FakeTaskRunner, PostTasksInPendingMode) { - FakeClock clock; - FakeTaskRunner task_runner{&clock, 1}; - task_runner.SetMode(FakeTaskRunner::Mode::kPending); - EXPECT_EQ(task_runner.GetConcurrentCount(), 1); - task_runner.PostTask([]() {}); - task_runner.PostTask([]() {}); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 2); - task_runner.RunNextPendingTask(); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 1); - task_runner.RunNextPendingTask(); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 0); -} - -TEST(FakeTaskRunner, RunAllPostedTasksInPendingMode) { - FakeClock clock; - FakeTaskRunner task_runner{&clock, 1}; - task_runner.SetMode(FakeTaskRunner::Mode::kPending); - task_runner.PostTask([]() {}); - task_runner.PostTask([]() {}); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 2); - task_runner.RunAllPendingTasks(); - ASSERT_TRUE( - FakeTaskRunner::WaitForRunningTasksWithTimeout(absl::Milliseconds(100))); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 0); -} - -TEST(FakeTaskRunner, PostDelayedTaskInPendingMode) { - FakeClock clock; - bool called = false; - FakeTaskRunner task_runner{&clock, 1}; - task_runner.SetMode(FakeTaskRunner::Mode::kPending); - task_runner.PostDelayedTask(absl::Seconds(1), [&called]() { called = true; }); - EXPECT_EQ(task_runner.GetAllDelayedTasks().size(), 1); - clock.FastForward(absl::Seconds(1)); - EXPECT_EQ(task_runner.GetAllDelayedTasks().size(), 0); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 1); - task_runner.RunAllPendingTasks(); - ASSERT_TRUE( - FakeTaskRunner::WaitForRunningTasksWithTimeout(absl::Milliseconds(100))); - EXPECT_EQ(task_runner.GetAllPendingTasks().size(), 0); - EXPECT_TRUE(called); } TEST(FakeTaskRunner, PostTasksRunInSequence) {