internal refactor

PiperOrigin-RevId: 547586065
This commit is contained in:
Guogang Li
2023-07-12 13:31:57 -07:00
committed by Copybara-Service
parent 0d2e75f2a2
commit 85a01e9636
3 changed files with 31 additions and 209 deletions
+17 -120
View File
@@ -14,8 +14,6 @@
#include "internal/test/fake_task_runner.h"
#include <atomic>
#include <future> // NOLINT
#include <memory>
#include <utility>
#include <vector>
@@ -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<void()> 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> timer = std::make_unique<FakeTimer>(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<absl::AnyInvocable<void()>>&
FakeTaskRunner::GetAllPendingTasks() const {
absl::MutexLock lock(&mutex_);
return pending_tasks_;
}
const absl::flat_hash_map<uint32_t, std::unique_ptr<Timer>>&
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<void()> task) {
CleanThreads();
++total_running_thread_count_;
// Run the task in a new thread, to simulate the real environment.
std::future<void> 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
+14 -41
View File
@@ -17,25 +17,24 @@
#include <atomic>
#include <cstdint>
#include <future> //NOLINT
#include <memory>
#include <vector>
#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<MultiThreadExecutor>(count)) {}
~FakeTaskRunner() override ABSL_LOCKS_EXCLUDED(mutex_);
bool PostTask(absl::AnyInvocable<void()> task) override
@@ -47,51 +46,25 @@ class FakeTaskRunner : public TaskRunner {
absl::AnyInvocable<void()> 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<absl::AnyInvocable<void()>>& GetAllPendingTasks() const
ABSL_LOCKS_EXCLUDED(mutex_);
const absl::flat_hash_map<uint32_t, std::unique_ptr<Timer>>&
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<void()> 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<MultiThreadExecutor> task_executor_ ABSL_GUARDED_BY(mutex_) =
nullptr;
// Used for pending mode
std::vector<absl::AnyInvocable<void()>> pending_tasks_
ABSL_GUARDED_BY(mutex_);
std::vector<absl::AnyInvocable<void()>> queued_tasks_ ABSL_GUARDED_BY(mutex_);
std::vector<uint32_t> completed_delayed_tasks_ ABSL_GUARDED_BY(mutex_);
absl::flat_hash_map<uint32_t, std::unique_ptr<Timer>> queued_delayed_tasks_
ABSL_GUARDED_BY(mutex_);
std::vector<std::future<void>> threads_ ABSL_GUARDED_BY(mutex_);
int running_thread_count_ ABSL_GUARDED_BY(mutex_) = 0;
// Tracks delayed tasks.
std::vector<std::unique_ptr<Timer>> timers_ ABSL_GUARDED_BY(mutex_);
static std::atomic_uint total_running_thread_count_;
};
-48
View File
@@ -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) {