diff --git a/internal/platform/task_runner_impl.cc b/internal/platform/task_runner_impl.cc index 587f5319..4f25101c 100644 --- a/internal/platform/task_runner_impl.cc +++ b/internal/platform/task_runner_impl.cc @@ -19,12 +19,17 @@ #include "absl/functional/any_invocable.h" #include "internal/platform/implementation/crypto.h" +#include "internal/platform/single_thread_executor.h" #include "internal/platform/timer_impl.h" namespace nearby { TaskRunnerImpl::TaskRunnerImpl(uint32_t runner_count) { - executor_ = std::make_unique<::nearby::MultiThreadExecutor>(runner_count); + if (runner_count == 1) { + executor_ = std::make_unique<::nearby::SingleThreadExecutor>(); + } else { + executor_ = std::make_unique<::nearby::MultiThreadExecutor>(runner_count); + } } TaskRunnerImpl::~TaskRunnerImpl() = default; diff --git a/internal/platform/task_runner_impl.h b/internal/platform/task_runner_impl.h index 3e1effb9..337a10df 100644 --- a/internal/platform/task_runner_impl.h +++ b/internal/platform/task_runner_impl.h @@ -50,7 +50,7 @@ class TaskRunnerImpl : public TaskRunner { uint64_t GenerateId(); mutable absl::Mutex mutex_; - std::unique_ptr<::nearby::MultiThreadExecutor> executor_; + std::unique_ptr<::nearby::SubmittableExecutor> executor_; absl::flat_hash_map> timers_map_ ABSL_GUARDED_BY(mutex_); }; diff --git a/internal/platform/task_runner_impl_test.cc b/internal/platform/task_runner_impl_test.cc index 4bc40a4e..b726a863 100644 --- a/internal/platform/task_runner_impl_test.cc +++ b/internal/platform/task_runner_impl_test.cc @@ -25,8 +25,12 @@ namespace nearby { namespace { -TEST(TaskRunnerImpl, PostTask) { - TaskRunnerImpl task_runner{1}; +constexpr uint32_t kNumThreads[] = {1, 10}; + +class BaseTaskRunnerImplTest : public ::testing::TestWithParam {}; + +TEST_P(BaseTaskRunnerImplTest, PostTask) { + TaskRunnerImpl task_runner{GetParam()}; absl::Notification notification; bool called = false; @@ -38,8 +42,8 @@ TEST(TaskRunnerImpl, PostTask) { EXPECT_TRUE(called); } -TEST(TaskRunnerImpl, PostSequenceTasks) { - TaskRunnerImpl task_runner{1}; +TEST_P(BaseTaskRunnerImplTest, PostSequenceTasks) { + TaskRunnerImpl task_runner{GetParam()}; std::vector completed_tasks; absl::Notification notification; @@ -66,8 +70,8 @@ TEST(TaskRunnerImpl, PostSequenceTasks) { EXPECT_EQ(completed_tasks[1], "task2"); } -TEST(TaskRunnerImpl, DISABLED_PostDelayedTask) { - TaskRunnerImpl task_runner{1}; +TEST_P(BaseTaskRunnerImplTest, DISABLED_PostDelayedTask) { + TaskRunnerImpl task_runner{GetParam()}; std::vector completed_tasks; absl::Notification notification; @@ -94,8 +98,8 @@ TEST(TaskRunnerImpl, DISABLED_PostDelayedTask) { EXPECT_EQ(completed_tasks[1], "task1"); } -TEST(TaskRunnerImpl, DISABLED_PostTwoDelayedTask) { - TaskRunnerImpl task_runner{1}; +TEST_P(BaseTaskRunnerImplTest, DISABLED_PostTwoDelayedTask) { + TaskRunnerImpl task_runner{GetParam()}; std::vector completed_tasks; absl::Notification notification; @@ -133,7 +137,26 @@ TEST(TaskRunnerImpl, DISABLED_PostTwoDelayedTask) { EXPECT_EQ(completed_tasks[2], "task3"); } -TEST(TaskRunnerImpl, PostTasksOnRunnerWithMultipleThreads) { +TEST(BaseTaskRunnerImplTest, PostTasksOnRunnerWithOneThread) { + TaskRunnerImpl task_runner{10}; + std::atomic_int count = 0; + absl::Notification notification; + + for (int i = 0; i < 10; i++) { + task_runner.PostTask([&count, ¬ification]() { + absl::SleepFor(absl::Milliseconds(100)); + count++; + if (count == 10) { + notification.Notify(); + } + }); + } + + notification.WaitForNotificationWithTimeout(absl::Milliseconds(1900)); + EXPECT_EQ(count, 10); +} + +TEST(BaseTaskRunnerImplTest, PostTasksOnRunnerWithMultipleThreads) { TaskRunnerImpl task_runner{10}; std::atomic_int count = 0; absl::Notification notification; @@ -152,11 +175,15 @@ TEST(TaskRunnerImpl, PostTasksOnRunnerWithMultipleThreads) { EXPECT_EQ(count, 10); } -TEST(TaskRunnerImpl, PostEmptyTask) { - TaskRunnerImpl task_runner{1}; +TEST_P(BaseTaskRunnerImplTest, PostEmptyTask) { + TaskRunnerImpl task_runner{GetParam()}; EXPECT_TRUE(task_runner.PostTask(nullptr)); EXPECT_TRUE(task_runner.PostDelayedTask(absl::Milliseconds(100), nullptr)); } +INSTANTIATE_TEST_SUITE_P(ParameterizedBasePcpHandlerTest, + BaseTaskRunnerImplTest, + ::testing::ValuesIn(kNumThreads)); + } // namespace } // namespace nearby