《C++ Concurrency In Action》并行快排示例是否存在线程安全Bug?
《C++ Concurrency In Action(第二版)》并行快排实现的数据竞争问题
问题确认
你指出的问题完全正确:这段代码中sorter类的threads成员是普通非线程安全的std::vector,存在明确的数据竞争风险。
当主线程在do_sort中执行threads.push_back(std::thread(&sorter<T>::sort_thread,this))时,新创建的线程会立刻进入sort_thread循环,调用try_sort_chunk处理任务块,而sort_chunk又会递归调用do_sort——如果此时threads.size() < max_thread_count的条件依然成立,新线程就会尝试向threads中push新元素,这就和主线程(或其他线程)的push_back操作形成无保护的并发写入,违反了C++标准中对std::vector的线程安全要求:只有当所有操作都是读取,或仅有一个线程进行写入时,std::vector才是安全的。
修复方案
方式1:用互斥锁保护threads的访问
为sorter类添加std::mutex成员,在所有访问threads的操作(读取size()、执行push_back()、析构时遍历join)中加锁:
template<typename T> struct sorter { // ... 原有成员 ... std::mutex threads_mutex; // 添加互斥锁 std::list<T> do_sort(std::list<T>& chunk_data) { // ... 原有代码 ... std::lock_guard<std::mutex> lock(threads_mutex); if(threads.size()<max_thread_count) { threads.push_back(std::thread(&sorter<T>::sort_thread,this)); } // ... 原有代码 ... } ~sorter() { end_of_data=true; std::lock_guard<std::mutex> lock(threads_mutex); for(unsigned i=0;i<threads.size();++i) { threads[i].join(); } } };
方式2:预先创建所有工作线程
既然max_thread_count是固定值,可在sorter构造函数中直接创建所有工作线程,彻底避免动态修改threads的并发风险:
sorter(): max_thread_count(std::thread::hardware_concurrency()-1), end_of_data(false) { for(unsigned i=0;i<max_thread_count;++i) { threads.push_back(std::thread(&sorter<T>::sort_thread,this)); } }
原代码清单
template<typename T> struct sorter { struct chunk_to_sort { std::list<T> data; std::promise<std::list<T>> promise; }; thread_safe_stack<chunk_to_sort> chunks; std::vector<std::thread> threads; unsigned const max_thread_count; std::atomic<bool> end_of_data; sorter(): max_thread_count(std::thread::hardware_concurrency()-1), end_of_data(false) {} ~sorter() { end_of_data=true; for(unsigned i=0;i<threads.size();++i) { threads[i].join(); } } void try_sort_chunk() { boost::shared_ptr<chunk_to_sort> chunk=chunks.pop(); if(chunk) { sort_chunk(chunk); } } std::list<T> do_sort(std::list<T>& chunk_data) { if(chunk_data.empty()) { return chunk_data; } std::list<T> result; result.splice(result.begin(),chunk_data,chunk_data.begin()); T const& partition_val=*result.begin(); typename std::list<T>::iterator divide_point= std::partition(chunk_data.begin(),chunk_data.end(), [&](T const& val){return val<partition_val;}); chunk_to_sort new_lower_chunk; new_lower_chunk.data.splice(new_lower_chunk.data.end(), chunk_data,chunk_data.begin(), divide_point); std::future<std::list<T>> new_lower= new_lower_chunk.promise.get_future(); chunks.push(std::move(new_lower_chunk)); if(threads.size()<max_thread_count) { threads.push_back(std::thread(&sorter<T>::sort_thread,this)); } std::list<T> new_higher(do_sort(chunk_data)); result.splice(result.end(),new_higher); while(new_lower.wait_for(std::chrono::seconds(0)) != std::future_status::ready) { try_sort_chunk(); } result.splice(result.begin(),new_lower.get()); return result; } void sort_chunk(boost::shared_ptr<chunk_to_sort> const& chunk) { chunk->promise.set_value(do_sort(chunk->data)); } void sort_thread() { while(!end_of_data) { try_sort_chunk(); std::this_thread::yield(); } } }; template<typename T> std::list<T> parallel_quick_sort(std::list<T> input) { if(input.empty()) { return input; } sorter<T> s; return s.do_sort(input); }
内容的提问来源于stack exchange,提问作者DwayneDuane
相关产品推荐
相关产品推荐

