diff --git a/src/bthread/bthread.cpp b/src/bthread/bthread.cpp index 727b4ebace..836dec9994 100644 --- a/src/bthread/bthread.cpp +++ b/src/bthread/bthread.cpp @@ -141,7 +141,10 @@ bthread_t init_for_pthread_stack_trace() { } pthread_fake_meta->attr = BTHREAD_ATTR_PTHREAD; - pthread_fake_meta->tid = make_tid(*pthread_fake_meta->version_butex, slot); + auto* version = reinterpret_cast*>( + pthread_fake_meta->version_butex); + pthread_fake_meta->tid = make_tid(static_cast( + version->load(butil::memory_order_relaxed)), slot); // Make TaskTracer use signal trace mode for pthread. c->_task_tracer.set_running_status(syscall(SYS_gettid), pthread_fake_meta); @@ -152,11 +155,17 @@ bthread_t init_for_pthread_stack_trace() { { BAIDU_SCOPED_LOCK(pthread_fake_meta->version_lock); tracing = TaskTracer::set_end_status_unsafe(pthread_fake_meta); - // If resulting version is 0, - // change it to 1 to make bthread_t never be 0. - if (0 == ++*pthread_fake_meta->version_butex) { - ++*pthread_fake_meta->version_butex; + // Publish the version atomically, just like task_runner(), since + // lock-free readers may still access this TaskMeta. + auto* version = reinterpret_cast*>( + pthread_fake_meta->version_butex); + uint32_t next_version = static_cast( + version->load(butil::memory_order_relaxed)) + 1; + // Skip zero to make bthread_t never be 0. + if (0 == next_version) { + ++next_version; } + version->store(static_cast(next_version), butil::memory_order_release); } if (tracing) { diff --git a/src/bthread/task_group.cpp b/src/bthread/task_group.cpp index 679e52ef5f..5b02f8c997 100644 --- a/src/bthread/task_group.cpp +++ b/src/bthread/task_group.cpp @@ -380,7 +380,9 @@ int TaskGroup::init(size_t runqueue_capacity) { m->cpuwide_start_ns = butil::cpuwide_time_ns(); m->stat = EMPTY_STAT; m->attr = BTHREAD_ATTR_TASKGROUP; - m->tid = make_tid(*m->version_butex, slot); + auto version = reinterpret_cast*>(m->version_butex); + m->tid = make_tid(static_cast( + version->load(butil::memory_order_relaxed)), slot); m->set_stack(stk); #ifdef BUTIL_USE_ASAN @@ -520,9 +522,17 @@ void TaskGroup::task_runner(intptr_t skip_remained) { #ifdef BRPC_BTHREAD_TRACER tracing = TaskTracer::set_end_status_unsafe(m); #endif // BRPC_BTHREAD_TRACER - if (0 == ++*m->version_butex) { - ++*m->version_butex; + // Bump the version with a release store so that it pairs with the + // acquire load in TaskGroup::join(): all memory writes made by this + // bthread become visible to the joining thread. Atomic access also + // avoids data races with the lock-free reads in join() and exists(). + auto* version = reinterpret_cast*>(m->version_butex); + uint32_t next_version = static_cast( + version->load(butil::memory_order_relaxed)) + 1; + if (0 == next_version) { + ++next_version; } + version->store(static_cast(next_version), butil::memory_order_release); } butex_wake_except(m->version_butex, 0); @@ -590,7 +600,9 @@ int TaskGroup::start_foreground(TaskGroup** pg, } m->cpuwide_start_ns = start_ns; m->stat = EMPTY_STAT; - m->tid = make_tid(*m->version_butex, slot); + auto version = reinterpret_cast*>(m->version_butex); + m->tid = make_tid(static_cast( + version->load(butil::memory_order_relaxed)), slot); TaskGroup* g = *pg; m->priority_index = g->_cur_meta->priority_index; @@ -662,7 +674,9 @@ int TaskGroup::start_background(bthread_t* __restrict th, } m->cpuwide_start_ns = start_ns; m->stat = EMPTY_STAT; - m->tid = make_tid(*m->version_butex, slot); + auto* version = reinterpret_cast*>(m->version_butex); + m->tid = make_tid(static_cast( + version->load(butil::memory_order_relaxed)), slot); m->priority_index = _cur_meta->priority_index; *th = m->tid; if (using_attr.flags & BTHREAD_LOG_START_AND_FINISH) { @@ -709,16 +723,18 @@ int TaskGroup::join(bthread_t tid, void** return_value) { return EINVAL; } const uint32_t expected_version = get_version(tid); - while (*m->version_butex == expected_version) { - if (butex_wait(m->version_butex, expected_version, nullptr) < 0 && + // Acquire load pairs with the release store performed when the joined + // bthread ends (see the version bump above), ensuring all of its memory + // writes are visible after join() returns. This matches the semantic + // guarantee provided by pthread_join() across supported architectures. + auto* version = reinterpret_cast*>(m->version_butex); + const int expected_version_int = static_cast(expected_version); + while (version->load(butil::memory_order_acquire) == expected_version_int) { + if (butex_wait(m->version_butex, expected_version_int, nullptr) < 0 && errno != EWOULDBLOCK && errno != EINTR) { return errno; } } - // Ensure all memory writes made by the joined bthread are visible to - // the joining thread after join returns. This matches the semantic - // guarantee provided by pthread_join() across supported architectures. - butil::atomic_thread_fence(butil::memory_order_acquire); if (return_value) { *return_value = nullptr; } @@ -729,7 +745,10 @@ bool TaskGroup::exists(bthread_t tid) { if (tid != 0) { // tid of bthread is never 0. TaskMeta* m = address_meta(tid); if (m != nullptr) { - return (*m->version_butex == get_version(tid)); + auto version = reinterpret_cast*>(m->version_butex); + // Only check liveness; unlike join(), no user data is acquired. + return static_cast(version->load(butil::memory_order_relaxed)) + == get_version(tid); } } return false; diff --git a/src/bthread/task_meta.h b/src/bthread/task_meta.h index 2dae2fea27..7ca52e04c8 100644 --- a/src/bthread/task_meta.h +++ b/src/bthread/task_meta.h @@ -85,10 +85,18 @@ struct TaskMeta { // Scheduling of the thread can be delayed. bool about_to_quit{false}; - // [Not Reset] guarantee visibility of version_butex. + // [Not Reset] Serializes the version bump at bthread end (in task_runner) + // with accessors that validate the version before touching other fields of + // this TaskMeta (get_attr/set_stopped/interrupt/set_butex_waiter/...). It + // makes their "check version then read/write field" sequence atomic w.r.t. + // the bump, so they never operate on a slot that got recycled in between. pthread_spinlock_t version_lock{}; - - // [Not Reset] only modified by one bthread at any time, no need to be atomic + + // [Not Reset] Backed by a butex (internally `butil::atomic`). The version + // bump at bthread end is published with a release store, and join() observes + // it with an acquire load so the joined bthread's prior writes are visible + // after join() returns. All lock-free accesses must be atomic; liveness + // checks and reads before publishing a new task only need relaxed loads. uint32_t* version_butex{nullptr}; // The identifier. It does not have to be here, however many code is diff --git a/test/bthread_unittest.cpp b/test/bthread_unittest.cpp index ce2b368c91..c86bc83e8a 100644 --- a/test/bthread_unittest.cpp +++ b/test/bthread_unittest.cpp @@ -26,6 +26,7 @@ #include #include "bthread/bthread.h" #include "bthread/unstable.h" +#include "bthread/task_group.h" #include "bthread/task_meta.h" #include "bvar/bvar.h" @@ -265,6 +266,75 @@ TEST_F(BthreadTest, bthread_join) { ASSERT_EQ(0, bthread_start_urgent(&th, nullptr, join_self, nullptr)); } +struct JoinVisibilityData { + int seed; + bool delay_write; + int values[64]; +}; + +void* write_join_visibility_data(void* arg) { + auto data = static_cast(arg); + if (data->delay_write) { + // Give the caller a chance to enter the join wait path before writing. + bthread_usleep(1000); + } + for (size_t i = 0; i < ARRAY_SIZE(data->values); ++i) { + data->values[i] = data->seed + static_cast(i); + } + return nullptr; +} + +void check_join_visibility(bool join_after_exit) { + for (int round = 0; round < 1000; ++round) { + JoinVisibilityData data = {}; + data.seed = round + 1; + data.delay_write = !join_after_exit; + bthread_t tid; + ASSERT_EQ(0, bthread_start_background( + &tid, nullptr, write_join_visibility_data, &data)); + if (join_after_exit) { + // exists() uses a relaxed load, so observing completion here does + // not acquire the worker's writes. join() must still do so even + // when it returns without waiting on the butex. + while (bthread::TaskGroup::exists(tid)) { + bthread_usleep(10); + } + } + ASSERT_EQ(0, bthread_join(tid, nullptr)); + // The payload is deliberately non-atomic and is read only after join. + // Do not add a lock or a release/acquire completion flag to this test: + // that would provide an alternative way to publish the worker's data. + for (size_t i = 0; i < ARRAY_SIZE(data.values); ++i) { + ASSERT_EQ(data.seed + static_cast(i), data.values[i]) + << "round=" << round << " index=" << i + << " join_after_exit=" << join_after_exit; + } + } +} + +void* join_visibility_caller(void* arg) { + const bool is_bthread = *static_cast(arg); + EXPECT_EQ(is_bthread, bthread_self() != 0); + check_join_visibility(false); + check_join_visibility(true); + return nullptr; +} + +TEST_F(BthreadTest, join_visibility_from_pthread) { + bool is_bthread = false; + pthread_t caller; + ASSERT_EQ(0, pthread_create(&caller, nullptr, join_visibility_caller, &is_bthread)); + ASSERT_EQ(0, pthread_join(caller, nullptr)); +} + +TEST_F(BthreadTest, join_visibility_from_bthread) { + bool is_bthread = true; + bthread_t caller; + ASSERT_EQ(0, bthread_start_background( + &caller, nullptr, join_visibility_caller, &is_bthread)); + ASSERT_EQ(0, bthread_join(caller, nullptr)); +} + void* change_errno(void* arg) { errno = (intptr_t)arg; return nullptr;