无锁单生产者单消费者(SPSC)队列元素顺序错乱问题
有界无锁SPSC队列元素顺序错乱问题修复
问题描述
实现了一个有界无锁单生产者单消费者(SPSC)队列,测试时发现元素插入顺序与取出顺序不一致:单线程向队列写入1亿个连续整数,另一个线程读取并检查连续性,运行时会输出类似Oh no! last value was: 219718 but this value is 220231的错误(每次数值不同)。错误出现的位置和writer_thread中的yield_frequency相关,值越小越容易在几百的数值处触发失败。
问题代码
spsc_queue.hpp
#pragma once #include <memory> #include <utility> #include <bit> #include <atomic> #include <optional> template <typename T> class ChannelReader; template <typename T> class ChannelWriter; // 无法直接创建此类实例,make_queue返回一对ChannelReader和ChannelWriter包装实例,防止多线程读写 template <typename T> class SpscQueue { friend ChannelReader<T>; friend ChannelWriter<T>; public: static std::pair<ChannelReader<T>, ChannelWriter<T>> make_queue(size_t size) { if (!std::__has_single_bit(size)) throw ""; std::shared_ptr<SpscQueue<T>> queue_ptr(new SpscQueue<T>(size)); return std::make_pair(ChannelReader<T>(queue_ptr), ChannelWriter<T>(queue_ptr)); } ~SpscQueue() { delete[] data; } private: SpscQueue (size_t _size): data(new T[_size]()), size(_size) {} alignas(64) T* data; const size_t size; alignas(64) std::atomic<size_t> read_idx{0}; alignas(64) std::atomic<size_t> write_idx{0}; }; // 允许单个线程读取SPSC队列的包装类 template <typename T> class ChannelReader { friend SpscQueue<T>; public: // 尝试从队列读取,成功返回值,失败返回std::nullopt std::optional<T> try_get_next() { size_t read_idx = queue->read_idx; size_t write_idx = cached_write_idx; if (write_idx <= read_idx) { write_idx = queue->write_idx; cached_write_idx = write_idx; } if (write_idx <= read_idx) { return std::nullopt; } else { queue->read_idx++; return std::move(queue->data[read_idx % queue->size]); } } ChannelReader(const ChannelReader&) = delete; ChannelReader& operator=(const ChannelReader&) = delete; ChannelReader(ChannelReader&&) = default; ChannelReader& operator=(ChannelReader&&) = default; private: ChannelReader(std::shared_ptr<SpscQueue<T>> ptr): queue(ptr) {} std::shared_ptr<SpscQueue<T>> queue; alignas(64) size_t cached_write_idx{0}; }; // 允许单个线程写入SPSC队列的包装类 template <typename T> class ChannelWriter { friend SpscQueue<T>; public: // 尝试向队列写入,成功返回true,失败返回false bool try_write_next(const T& obj) { size_t read_idx = cached_read_idx; size_t write_idx = queue->write_idx; if (write_idx >= read_idx + queue->size) { read_idx = queue->read_idx; cached_read_idx = read_idx; } if (write_idx >= read_idx + queue->size) { return false; } else { queue->data[write_idx % queue->size] = obj; ++queue->write_idx; return true; } } ChannelWriter(const ChannelWriter&) = delete; ChannelWriter& operator=(const ChannelWriter&) = delete; ChannelWriter(ChannelWriter&&) = default; ChannelWriter& operator=(ChannelWriter&&) = default; private: ChannelWriter(std::shared_ptr<SpscQueue<T>> ptr): queue(ptr) {} std::shared_ptr<SpscQueue<T>> queue; alignas(64) size_t cached_read_idx{0}; };
main.cpp
#include <atomic> #include <chrono> #include <iostream> #include <thread> #include "../include/spsc_queue.hpp" namespace chrono = std::chrono; int main() { std::atomic<bool> latch{false}; auto [reader, writer] = SpscQueue<int>::make_queue(512); std::thread reader_thread([_reader = std::move(reader), &latch]() mutable { latch.wait(false); std::cerr << "Starting reader...\n"; int last_val = -1; while (last_val != 100'000'000 - 1) { if (auto data = _reader.try_get_next()) { if (*data != last_val + 1) { std::cerr << "Oh no! last value was: " << last_val << " but this value is " << *data << '\n'; std::exit(1); } last_val = *data; // std::cerr << last_val << " Read\n"; } } }); std::thread writer_thread([_writer = std::move(writer), &latch]() mutable { latch.wait(false); std::cerr << "Starting writer...\n"; for (int i = 0; i < 100'000'000; ++i) { for (int j = 0; !_writer.try_write_next(i); ++j) { constexpr int yield_frequency = 1 << 0; if (j % yield_frequency) std::this_thread::yield(); } // std::cerr << "writer wrote value" << i << '\n'; // if (i == 10'000'000) std::exit(1); } }); std::cout << "Start" << std::endl; { auto start = chrono::steady_clock::now(); latch = true; latch.notify_all(); reader_thread.join(); writer_thread.join(); auto finish = chrono::steady_clock::now(); auto elapsed_seconds = chrono::duration_cast<chrono::duration<double>>(finish - start).count(); std::cout << elapsed_seconds << std::endl; } std::cout << "End" << std::endl; }
问题根源与修复方案
核心问题
- 操作顺序颠倒:
try_get_next中先递增read_idx再读取数据,会导致写入线程认为该位置已被读取并覆盖数据,读线程最终读到覆盖后的新值,破坏顺序。 - 内存可见性未明确:原子变量的读写未指定内存顺序,可能导致指令重排或数据同步延迟。
修复步骤
1. 修正try_get_next操作逻辑
先读取数据,再更新read_idx,同时指定内存顺序保证可见性:
std::optional<T> try_get_next() { size_t current_read = queue->read_idx.load(std::memory_order_acquire); size_t current_write = cached_write_idx; if (current_write <= current_read) { current_write = queue->write_idx.load(std::memory_order_acquire); cached_write_idx = current_write; } if (current_write <= current_read) { return std::nullopt; } // 先读数据,再更新索引 T value = std::move(queue->data[current_read % queue->size]); queue->read_idx.store(current_read + 1, std::memory_order_release); return value; }
2. 修正try_write_next操作逻辑
先写入数据,再更新write_idx,补充内存顺序:
bool try_write_next(const T& obj) { size_t current_write = queue->write_idx.load(std::memory_order_acquire); size_t current_read = cached_read_idx; if (current_write >= current_read + queue->size) { current_read = queue->read_idx.load(std::memory_order_acquire); cached_read_idx = current_read; if (current_write >= current_read + queue->size) { return false; } } // 先写数据,再更新索引 queue->data[current_write % queue->size] = obj; queue->write_idx.store(current_write + 1, std::memory_order_release); return true; }
3. 优化异常提示
将构造队列时的空异常改为有意义的提示:
if (!std::__has_single_bit(size)) throw "Queue size must be a power of two";
修复后完整spsc_queue.hpp
#pragma once #include <memory> #include <utility> #include <bit> #include <atomic> #include <optional> template <typename T> class ChannelReader; template <typename T> class ChannelWriter; // 无法直接创建此类实例,make_queue返回一对ChannelReader和ChannelWriter包装实例,防止多线程读写 template <typename T> class SpscQueue { friend ChannelReader<T>; friend ChannelWriter<T>; public: static std::pair<ChannelReader<T>, ChannelWriter<T>> make_queue(size_t size) { if (!std::__has_single_bit(size)) throw "Queue size must be a power of two"; std::shared_ptr<SpscQueue<T>> queue_ptr(new SpscQueue<T>(size)); return std::make_pair(ChannelReader<T>(queue_ptr), ChannelWriter<T>(queue_ptr)); } ~SpscQueue() { delete[] data; } private: SpscQueue (size_t _size): data(new T[_size]()), size(_size) {} alignas(64) T* data; const size_t size; alignas(64) std::atomic<size_t> read_idx{0}; alignas(64) std::atomic<size_t> write_idx{0}; }; // 允许单个线程读取SPSC队列的包装类 template <typename T> class ChannelReader { friend SpscQueue<T>; public: // 尝试从队列读取,成功返回值,失败返回std::nullopt std::optional<T> try_get_next() { size_t current_read = queue->read_idx.load(std::memory_order_acquire); size_t current_write = cached_write_idx; if (current_write <= current_read) { current_write = queue->write_idx.load(std::memory_order_acquire); cached_write_idx = current_write; } if (current_write <= current_read) { return std::nullopt; } T value = std::move(queue->data[current_read % queue->size]); queue->read_idx.store(current_read + 1, std::memory_order_release); return value; } ChannelReader(const ChannelReader&) = delete; ChannelReader& operator=(const ChannelReader&) = delete; ChannelReader(ChannelReader&&) = default; ChannelReader& operator=(ChannelReader&&) = default; private: ChannelReader(std::shared_ptr<SpscQueue<T>> ptr): queue(ptr) {} std::shared_ptr<SpscQueue<T>> queue; alignas(64) size_t cached_write_idx{0}; }; // 允许单个线程写入SPSC队列的包装类 template <typename T> class ChannelWriter { friend SpscQueue<T>; public: // 尝试向队列写入,成功返回true,失败返回false bool try_write_next(const T& obj) { size_t current_write = queue->write_idx.load(std::memory_order_acquire); size_t current_read = cached_read_idx; if (current_write >= current_read + queue->size) { current_read = queue->read_idx.load(std::memory_order_acquire); cached_read_idx = current_read; if (current_write >= current_read + queue->size) { return false; } } queue->data[current_write % queue->size] = obj; queue->write_idx.store(current_write + 1, std::memory_order_release); return true; } ChannelWriter(const ChannelWriter&) = delete; ChannelWriter& operator=(const ChannelWriter&) = delete; ChannelWriter(ChannelWriter&&) = default; ChannelWriter& operator=(ChannelWriter&&) = default; private: ChannelWriter(std::shared_ptr<SpscQueue<T>> ptr): queue(ptr) {} std::shared_ptr<SpscQueue<T>> queue; alignas(64) size_t cached_read_idx{0}; };
修复原理
- 操作顺序修正:确保数据读写操作在索引更新前完成,避免数据被提前覆盖。
- 内存顺序保障:
memory_order_acquire保证后续操作能看到原子变量的最新值;memory_order_release保证之前的写入操作对其他线程可见,同时兼顾性能。 - 缓存逻辑优化:仅在缓存失效时重新加载原子变量,减少不必要的原子操作开销。
内容的提问来源于stack exchange,提问作者ayaan098
相关产品推荐
相关产品推荐

