diff --git a/cyber/timer/timer.h b/cyber/timer/timer.h index 99d6cd74b1c..54dc2e18b7c 100644 --- a/cyber/timer/timer.h +++ b/cyber/timer/timer.h @@ -25,6 +25,10 @@ namespace apollo { namespace cyber { +namespace timer { +class TimerConcurrencyTest; +} // namespace timer + /** * @brief The options of timer * @@ -110,6 +114,8 @@ class Timer { void Stop(); private: + friend class timer::TimerConcurrencyTest; + bool InitTimerTask(); uint64_t timer_id_; TimerOption timer_opt_; diff --git a/cyber/timer/timer_test.cc b/cyber/timer/timer_test.cc index 8eb6f93be9a..fa1ae8e332b 100644 --- a/cyber/timer/timer_test.cc +++ b/cyber/timer/timer_test.cc @@ -18,14 +18,21 @@ #include "cyber/timer/timer.h" +#include +#include +#include +#include #include +#include #include +#include #include "gtest/gtest.h" #include "cyber/common/util.h" #include "cyber/cyber.h" #include "cyber/init.h" +#include "cyber/scheduler/scheduler_factory.h" namespace apollo { namespace cyber { @@ -34,6 +41,138 @@ namespace timer { using cyber::Timer; using cyber::TimerOption; +namespace { + +bool WaitForCount(const std::atomic& count, uint32_t expected) { + const auto deadline = + std::chrono::steady_clock::now() + std::chrono::seconds(5); + while (std::chrono::steady_clock::now() < deadline) { + if (count.load() == expected) { + return true; + } + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + return count.load() == expected; +} + +} // namespace + +class TimerConcurrencyTest : public testing::Test { + protected: + void SetUp() override { + timing_wheel_ = TimingWheel::Instance(); + timing_wheel_->Shutdown(); + ASSERT_FALSE(timing_wheel_->tick_thread_.joinable()); + timing_wheel_->running_ = true; + } + + void TearDown() override { timing_wheel_->Shutdown(); } + + bool InitializeTimer(Timer* timer) { + if (!timer->InitTimerTask()) { + return false; + } + timer->started_.store(true); + return true; + } + + std::weak_ptr QueueTimerCallback(Timer* timer) { + auto task = timer->task_; + auto& bucket = + timing_wheel_->work_wheel_[timing_wheel_->current_work_wheel_index_]; + bucket.AddTask(task); + timing_wheel_->Tick(); + return task; + } + + private: + TimingWheel* timing_wheel_ = nullptr; +}; + +TEST_F(TimerConcurrencyTest, stop_before_queued_callback_runs) { + const uint32_t task_pool_size = scheduler::Instance()->TaskPoolSize(); + ASSERT_GT(task_pool_size, 0); + + std::atomic callback_count{0}; + auto callback_state = std::make_shared(0); + std::weak_ptr callback_lifetime = callback_state; + auto timer = std::make_unique( + TIMER_RESOLUTION_MS, + [callback_state, &callback_count] { + ++(*callback_state); + callback_count.fetch_add(1); + }, + true); + callback_state.reset(); + ASSERT_TRUE(InitializeTimer(timer.get())); + + // Occupy every TaskManager consumer so Tick() can only queue the callback. + std::promise release_blockers_promise; + auto release_blockers = release_blockers_promise.get_future().share(); + std::atomic blockers_started{0}; + std::vector> blockers; + blockers.reserve(task_pool_size); + for (uint32_t i = 0; i < task_pool_size; ++i) { + blockers.emplace_back(Async([&blockers_started, release_blockers] { + blockers_started.fetch_add(1); + release_blockers.wait(); + })); + } + + if (!WaitForCount(blockers_started, task_pool_size)) { + release_blockers_promise.set_value(); + for (auto& blocker : blockers) { + blocker.wait(); + } + FAIL() << "Failed to occupy the task pool"; + return; + } + + auto task_lifetime = QueueTimerCallback(timer.get()); + timer->Stop(); + timer.reset(); + // Queued work must not retain the stopped task or callback-owned state. + EXPECT_TRUE(task_lifetime.expired()); + EXPECT_TRUE(callback_lifetime.expired()); + + // A full second wave is a queue drain fence: its last task cannot start + // until the earlier timer callback has finished. + std::promise release_drains_promise; + auto release_drains = release_drains_promise.get_future().share(); + std::atomic drains_started{0}; + std::vector> drains; + drains.reserve(task_pool_size); + for (uint32_t i = 0; i < task_pool_size; ++i) { + drains.emplace_back(Async([&drains_started, release_drains] { + drains_started.fetch_add(1); + release_drains.wait(); + })); + } + + release_blockers_promise.set_value(); + if (!WaitForCount(drains_started, task_pool_size)) { + release_drains_promise.set_value(); + for (auto& blocker : blockers) { + blocker.wait(); + } + for (auto& drain : drains) { + drain.wait(); + } + FAIL() << "Failed to drain the queued timer callback"; + return; + } + + release_drains_promise.set_value(); + for (auto& blocker : blockers) { + blocker.get(); + } + for (auto& drain : drains) { + drain.get(); + } + + EXPECT_EQ(callback_count.load(), 0); +} + TEST(TimerTest, one_shot) { int count = 0; Timer timer( @@ -79,6 +218,23 @@ TEST(TimerTest, start_stop) { } } +TEST(TimerTest, periodic_callback_preserves_state) { + std::promise second_fire_promise; + auto second_fire = second_fire_promise.get_future(); + Timer timer( + TIMER_RESOLUTION_MS, + [count = 0, &second_fire_promise]() mutable { + if (++count == 2) { + second_fire_promise.set_value(); + } + }, + false); + timer.Start(); + EXPECT_EQ(second_fire.wait_for(std::chrono::seconds(1)), + std::future_status::ready); + timer.Stop(); +} + TEST(TimerTest, sim_mode) { auto count = 0; diff --git a/cyber/timer/timing_wheel.cc b/cyber/timer/timing_wheel.cc index 179e8c739de..5ac868dd60f 100644 --- a/cyber/timer/timing_wheel.cc +++ b/cyber/timer/timing_wheel.cc @@ -53,11 +53,13 @@ void TimingWheel::Tick() { if (task) { ADEBUG << "index: " << current_work_wheel_index_ << " timer id: " << task->timer_id_; - auto* callback = - reinterpret_cast*>(&(task->callback)); - cyber::Async([this, callback] { + std::weak_ptr task_weak_ptr = task; + cyber::Async([this, task_weak_ptr] { if (this->running_) { - (*callback)(); + auto task = task_weak_ptr.lock(); + if (task) { + task->callback(); + } } }); } diff --git a/cyber/timer/timing_wheel.h b/cyber/timer/timing_wheel.h index 92dde7c4cbf..ab3af768201 100644 --- a/cyber/timer/timing_wheel.h +++ b/cyber/timer/timing_wheel.h @@ -31,6 +31,10 @@ namespace apollo { namespace cyber { +namespace timer { +class TimerConcurrencyTest; +} // namespace timer + struct TimerTask; static const uint64_t WORK_WHEEL_SIZE = 512; @@ -65,6 +69,8 @@ class TimingWheel { inline uint64_t TickCount() const { return tick_count_; } private: + friend class timer::TimerConcurrencyTest; + inline uint64_t GetWorkWheelIndex(const uint64_t index) { return index & (WORK_WHEEL_SIZE - 1); }