Create archive for vibe code

This commit is contained in:
12hydrogen
2026-06-25 17:37:39 +08:00
parent 432238635b
commit 973b2e60de
5 changed files with 1043 additions and 65 deletions
+13 -6
View File
@@ -2,10 +2,10 @@ cmake_minimum_required(VERSION 3.10)
project(WinSock2Wrapper)
set(C_STANDARD 17)
set(CXX_STANDARD 17)
set(C_STANDARD_REQUIRED True)
set(CXX_STANDARD_REQUIRED True)
set(CMAKE_C_STANDARD 17)
set(CMAKE_CXX_STANDARD 17)
set(CMAKE_C_STANDARD_REQUIRED True)
set(CMAKE_CXX_STANDARD_REQUIRED True)
add_library(
wrapper
@@ -32,7 +32,14 @@ target_link_libraries(
PUBLIC wrapper
)
# ThreadPool + LockFreeQueue test (header-only, no link deps)
add_executable(
test_shift
./shift.cpp
test_pool
./test_pool.cpp
./pool.cpp
)
target_include_directories(
test_pool
PUBLIC .
)
+6 -7
View File
@@ -1,8 +1,7 @@
// pool.cpp — ThreadPool & LockFreeQueue implementation file
//
// Note: All template implementations are in pool.hpp (header-only).
// This file exists for future non-template extensions or
// explicit template instantiations if needed.
#include "pool.hpp"
int main() {
Queue<int> queue;
queue.push(1);
auto load = queue.pop();
}
+663 -30
View File
@@ -1,44 +1,677 @@
#include <thread>
#pragma once
/**
* @file pool.hpp
* @brief 无锁有界 MPMC 队列 + 固定大小线程池,适用于多线程网络通信场景。
*
* 本文件提供两个组件:
* 1. LockFreeQueue<T, N> — Vyukov 有界多生产者多消费者无锁队列
* 2. ThreadPool — 基于 LockFreeQueue 的固定线程数线程池
*
* 设计原则:
* - 纯 C++17 标准库,无第三方依赖
* - 有界队列提供天然背压(backpressure),防止过载时 OOM
* - 混合等待策略:自旋(低延迟)+ 条件变量(低 CPU)
* - 详细注释解释「做什么」与「为什么这样做」
*
* 参考:
* - Vyukov bounded MPMC queue:
* http://www.1024cores.net/home/lock-free-algorithms/queues/bounded-mpmc-queue
*/
#include <mutex>
#include <future>
#include <atomic>
#include <chrono>
#include <condition_variable>
#include <memory>
#include <cstddef>
#include <cstdint>
#include <functional>
#include <future>
#include <mutex>
#include <thread>
#include <type_traits>
template <typename T>
struct QNode {
std::unique_ptr<T> payload;
volatile std::atomic<QNode*> next;
#include <vector>
QNode() = default;
// ============================================================================
// LockFreeQueue — Vyukov 有界 MPMC 无锁队列
// ============================================================================
explicit QNode(T&& payload):
payload(std::forward<T>(payload)) {}
/**
* @brief 有界、无锁、多生产者多消费者(MPMC)队列。
*
* 算法来源:Dmitry Vyukov 的有界 MPMC 队列
* 参考:http://www.1024cores.net/home/lock-free-algorithms/queues/bounded-mpmc-queue
*
* ## 为什么选择这个算法(而非 Michael-Scott 链表队列)?
*
* 1. **无内存分配**:数组预分配,enqueue/dequeue 不涉及 new/delete。
* 网络服务器中每秒可能有数万次入队/出队,避免每次分配内存可显著减少延迟抖动。
*
* 2. **缓存友好**:环形数组内存连续,CPU 预取效果好。
* 链表节点分散在堆上,每次访问都可能触发 cache miss。
*
* 3. **无 ABA 问题**:每个 cell 有单调递增的 sequence 号,永不复用。
* Michael-Scott 队列需要 hazard pointer / epoch-based reclamation 等额外机制来解决 ABA
* 本算法天然免疫。
*
* 4. **真正的 MPMC**:多个生产者和多个消费者可以并发操作不同 cell,
* 互不阻塞。吞吐量随线程数线性扩展。
*
* ## 为什么选择有界队列?
*
* - 网络服务器的处理能力有限(CPU 核数 × 每任务耗时)
* - 有界队列在满时返回 false → 调用方可以感知过载并做限流/降级
* - 无界队列在流量尖峰时会无限增长 → OOM
* - 有界 = 天然背压(backpressure
*
* @tparam T 元素类型(必须可移动;通过移动语义避免拷贝)
* @tparam N 容量 — 必须为 2 的幂(bitwise 取模要求)
*/
template <typename T, size_t N>
class LockFreeQueue {
// ========================================================================
// 编译期校验
// ========================================================================
// N 必须是 2 的幂。
// 原因:取模运算 `pos % N` 在 N 为 2 的幂时可以优化为 `pos & (N-1)`
// 这是一个单周期位运算,比整数除法快 10-50 倍。在热路径上至关重要。
static_assert((N & (N - 1)) == 0,
"LockFreeQueue: N must be a power of 2. "
"This allows fast modulo via bitwise AND: pos & (N-1)");
// 容量至少为 2。
// 当 N=1 时,head 和 tail 指向同一个 cell,读写互相覆盖,算法失效。
static_assert(N >= 2,
"LockFreeQueue: capacity must be at least 2");
// ========================================================================
// Cell:环形缓冲区中的一个槽位
// ========================================================================
/**
* 每个 cell 包含一个原子 sequence 号和元素数据。
*
* ## sequence 的作用(核心设计)
*
* 为什么每个 cell 需要独立的 sequence,而不是只靠全局 head/tail
*
* 在 MPMC 场景下,多个线程可能同时访问同一个 cell(例如 tail 刚推进,
* 但上一个 writer 还没写完)。sequence 充当每个 cell 的「回合指示器」:
*
* - `sequence == position` → cell 空闲,等待写入
* - `sequence == position + 1` → cell 已写入,等待读取
* - `sequence == position + N` → cell 已读取,等待下一轮写入
*
* 由于 sequence 单调递增、永不复用,即使 head/tail 绕过整个环回到
* 同一 cellsequence 的值也不同,CAS 永远不会在错误回合成功。
* 这就是「天然免疫 ABA」的原因。
*/
struct Cell {
std::atomic<size_t> sequence;
T data;
Cell() noexcept : sequence(0) {}
};
template <typename T>
class Queue {
private:
volatile std::atomic<QNode<T>*> head;
volatile std::atomic<QNode<T>*> tail;
// ========================================================================
// 环形缓冲区
// ========================================================================
Cell buffer_[N];
// ========================================================================
// head 和 tail 计数器(缓存行对齐,避免伪共享)
// ========================================================================
/**
* 为什么使用 alignas(64)
*
* x86 典型缓存行大小为 64 字节。如果不加对齐:
* - head 和 tail 可能落在同一缓存行
* - producer 更新 tail 时,consumer 核心上的缓存行被 invalidate
* - consumer 更新 head 时,producer 核心上的缓存行被 invalidate
* - 即使它们操作的是不同变量,也互相刷缓存 → 伪共享 → 性能剧降
*
* alignas(64) 确保 head 和 tail 在不同缓存行,消除伪共享。
* 这是高性能无锁数据结构的关键细节。
*/
alignas(64) std::atomic<size_t> head_;
alignas(64) std::atomic<size_t> tail_;
public:
void push(T&& payload) {
QNode<T>* node = new QNode(std::forward<T>(payload));
node->next = node;
// ========================================================================
// 构造
// ========================================================================
QNode<T>* _tail = tail;
QNode<T>* _next = _tail->next;
do {
node->next = _next;
} while (!_tail->next.compare_exchange_weak(_next, node));
tail.compare_exchange_strong(_tail, node);
/**
* 初始化所有 cellsequence[i] = i。
*
* 为什么 sequence[i] = i
*
* 构造时,head = 0, tail = 0。
* cell[0].sequence = 0 → 对于 pos=0seq == pos(空闲可写)
* cell[1].sequence = 1 → 对于 pos=1seq == pos(空闲可写)
* ...
*
* 队列初始为空:dequeue 检查 seq == pos+1(即 seq == 1),
* 但 seq 实际为 0(或更大),不相等 → 正确返回"空"。
*
* memory_order_relaxed 在此足够:
* - 构造是单线程的(无其他线程持有引用)
* - C++ 内存模型保证:构造线程的写入在线程创建后对子线程可见
* happens-before 通过 thread 构造函数传递)
*/
LockFreeQueue() noexcept
: head_(0)
, tail_(0)
{
for (size_t i = 0; i < N; ++i) {
buffer_[i].sequence.store(i, std::memory_order_relaxed);
}
}
std::unique_ptr<T> pop() {
// ========================================================================
// 禁止拷贝/移动(原子成员不可拷贝)
// ========================================================================
LockFreeQueue(const LockFreeQueue&) = delete;
LockFreeQueue& operator=(const LockFreeQueue&) = delete;
LockFreeQueue(LockFreeQueue&&) = delete;
LockFreeQueue& operator=(LockFreeQueue&&) = delete;
// ========================================================================
// enqueue:生产者入队
// ========================================================================
/**
* 尝试入队一个元素。非阻塞,立即返回。
*
* @param data 要入队的元素(移动语义,不拷贝)
* @return true 入队成功,false 队列已满
*
* ## 算法逐步解析
*
* 1. 用 relaxed 读取 tail 位置(精确位置由后续 CAS 保证)
* 2. 用 acquire 读取 cell 的 sequence(需要看到上一轮 writer 的数据)
* 3. 计算 signed 差值 `dif = seq - pos`
*
* | dif | 含义 | 动作 |
* |------|------|------|
* | == 0 | cell 空闲,可写入 | CAS(tail): 成功→写入+发布; 失败→重试 |
* | < 0 | 队列已满 | 返回 false |
* | > 0 | 其他 producer 已写入此 cell | 重新加载 tail 重试 |
*
* ## 为什么对 tail 做 CAS 而不是对 sequence 做 CAS
*
* tail 决定「写哪个 cell」。CAS tail 保证了只有一个 producer 能
* 获得当前 cell 的写入权。sequence 的更新则负责与 consumer 同步。
* 两个 CAS 各自分工:tail CAS 做互斥,sequence store 做发布。
*
* ## 内存序说明
*
* - sequence.load(acquire):确保看到上一轮 writer 写入的数据
* - tail.CAS(relaxed):不需要在此建立顺序关系,数据发布由下面的
* sequence.store(release) 完成
* - sequence.store(release):使写入的数据对 consumer 的 acquire 可见
* synchronizes-with 关系)
*/
bool enqueue(T&& data) noexcept {
// relaxed:精确值由后续 CAS 保证
size_t pos = tail_.load(std::memory_order_relaxed);
for (;;) {
// & (N-1):位运算取模,因为 N 是 2 的幂
Cell* cell = &buffer_[pos & (N - 1)];
// acquire:需要看到上一轮操作写入的数据
size_t seq = cell->sequence.load(std::memory_order_acquire);
// 有符号差值正确处理了无符号整数的环绕
// 例:seq=2, pos=6, N=4: dif = 2-6 = -4 < 0 → 队列满(正确)
intptr_t dif = static_cast<intptr_t>(seq) - static_cast<intptr_t>(pos);
if (dif == 0) {
// cell 空闲,尝试认领此位置
if (tail_.compare_exchange_weak(pos, pos + 1,
std::memory_order_relaxed)) {
// 认领成功!写入数据并发布。
cell->data = std::move(data);
// release:将数据对 consumer 可见。
// consumer 通过 acquire load 看到 pos+1 即知数据就绪。
cell->sequence.store(pos + 1, std::memory_order_release);
return true;
}
// CAS 失败:另一个 producer 抢先认领了此位置。
// compare_exchange_weak 已将 pos 更新为当前 tail,继续循环。
}
else if (dif < 0) {
// 队列满。
// seq < pos 表示 cell 的 sequence 还停留在上一轮,
// consumer 尚未通过 store(pos+N) 释放此 cell。
return false;
}
else {
// dif > 0:另一个 producer 已经认领并写入此 cell。
// 重新加载 tail,跳到下一个可用位置。
pos = tail_.load(std::memory_order_relaxed);
}
}
}
// ========================================================================
// dequeue:消费者出队
// ========================================================================
/**
* 尝试出队一个元素。非阻塞,立即返回。
*
* @param data 输出参数:接收出队的元素
* @return true 出队成功,false 队列为空
*
* ## 算法逐步解析
*
* 1. 用 relaxed 读取 head 位置
* 2. 用 acquire 读取 cell 的 sequence(需要看到 producer 写入的数据)
* 3. 计算 dif = seq - (pos + 1)
*
* | dif | 含义 | 动作 |
* |------|------|------|
* | == 0 | data 就绪,可读取 | CAS(head): 成功→读取+释放cell; 失败→重试 |
* | < 0 | 队列为空 | 返回 false |
* | > 0 | 其他 consumer 已取走此 cell | 重新加载 head 重试 |
*
* ## 出队后为什么 store(pos + N)
*
* 经过 N 次 enqueue 后,tail 将到达 pos+N 并尝试写入此 cell。
* 此时 sequence 为 pos+N 刚好等于新 tail → 判断为空闲可写。
* 这样环形缓冲区的复用就完成了:同一 cell 服务 0, N, 2N, ... 号回合。
*/
bool dequeue(T& data) noexcept {
size_t pos = head_.load(std::memory_order_relaxed);
for (;;) {
Cell* cell = &buffer_[pos & (N - 1)];
// acquire:需要看到 producer 通过 release store 发布的数据
size_t seq = cell->sequence.load(std::memory_order_acquire);
// 数据就绪的标志是 seq == pos + 1
// 注:seq == pos 对应 cell 空闲(enqueue 之前的状态)→ 队列空
intptr_t dif = static_cast<intptr_t>(seq) - static_cast<intptr_t>(pos + 1);
if (dif == 0) {
// 数据就绪,尝试认领此位置
if (head_.compare_exchange_weak(pos, pos + 1,
std::memory_order_relaxed)) {
// 认领成功!取出数据并释放 cell。
data = std::move(cell->data);
// release:告知未来 writer(在位置 pos+N)此 cell 已可用。
// writer 通过 acquire load 看到 pos+N 即知可写入。
cell->sequence.store(pos + N, std::memory_order_release);
return true;
}
// CAS 失败:另一个 consumer 抢先了。
// pos 已更新为当前 head,继续循环。
}
else if (dif < 0) {
// 队列空。
// seq == pos 表示 producer 还没有写入此 cell。
return false;
}
else {
// dif > 0:另一个 consumer 已经认领并释放此 cell。
// 重新加载 head。
pos = head_.load(std::memory_order_relaxed);
}
}
}
// ========================================================================
// 状态查询(best-effort,返回值可能立刻过时)
// ========================================================================
/**
* @return 队列是否(大概)为空
*
* relaxed 即可:这些查询本身是近似的,
* 返回值可能在调用者拿到之前就已过时。
*/
bool empty() const noexcept {
return head_.load(std::memory_order_relaxed) ==
tail_.load(std::memory_order_relaxed);
}
/// @return 队列是否(大概)已满
bool full() const noexcept {
size_t h = head_.load(std::memory_order_relaxed);
size_t t = tail_.load(std::memory_order_relaxed);
return (t - h) >= N;
}
/// @return 队列中大约有多少元素
size_t size() const noexcept {
size_t h = head_.load(std::memory_order_relaxed);
size_t t = tail_.load(std::memory_order_relaxed);
return (t >= h) ? (t - h) : 0;
}
};
// ============================================================================
// ThreadPool — 固定大小线程池
// ============================================================================
/**
* @brief 固定线程数的简单线程池,使用无锁任务队列。
*
* ## 适用场景
*
* - 多线程网络服务器(accept / read / write / process
* - 任务粒度小、数量大(每秒数千到数百万个任务)
* - 对延迟敏感(无锁队列避免 mutex 竞争)
* - 并发度可预测(固定 worker 数,不动态伸缩)
*
* ## 关键设计决策
*
* ### 1. 固定线程数(不动态伸缩)
* 线程创建/销毁开销大(内核态切换、栈分配),会引入延迟抖动。
* 网络服务器的并发需求通常可预测 → 固定线程 = 可预测性能。
* 默认值:`std::thread::hardware_concurrency()`。
*
* ### 2. 混合等待策略(自旋 + 条件变量)
* ```
* 快速路径:自旋尝试无锁出队 → 微秒级延迟(有任务时)
* 慢速路径:condition_variable 等待 → 零 CPU(空闲时)
* ```
* - 纯自旋:CPU 占用高,浪费电力,影响同机其他服务
* - 纯 mutex:高负载下 mutex 本身成为瓶颈
* - 混合:取两者之长,实际负载下最优
*
* ### 3. 有界队列 = 天然背压
* submit() 在队列满时自旋等待 → 调用方感知到过载。
* 避免无界队列导致的 OOM 风险。
*
* ### 4. std::packaged_task + std::future 返回值
* submit() 返回 future,调用方可获取任务执行结果。
* 对于网络 I/O 常见的回调模式非常实用。
*
* ## 生命周期
*
* ```
* 构造 → spawn workers → [submit...] → 析构 → 排空队列 → join workers
* ```
*
* 析构时保证所有已提交任务执行完毕(不会丢任务)。
*/
class ThreadPool {
public:
// ========================================================================
// 构造 / 析构
// ========================================================================
/**
* @param num_threads worker 线程数。
* 默认值:hardware_concurrency()。
* 最小值:1(强制)。
*
* 为什么不允许 0 个线程?
* 0 线程的池会静默丢弃所有任务 → 永不会是期望行为。
*/
explicit ThreadPool(size_t num_threads = std::thread::hardware_concurrency())
: done_(false)
{
if (num_threads == 0) {
num_threads = 1;
}
// 预分配 vector 空间,避免线程创建过程中的 reallocation
// (虽然 realloc 不常见,但预分配更安全、更可预测)
workers_.reserve(num_threads);
for (size_t i = 0; i < num_threads; ++i) {
// 每个 worker 执行 worker_loop,传入 this 指针。
//
// 使用 this 指针是安全的,因为:
// - ThreadPool 的生命周期长于其 worker 线程
// (析构函数先设置 done_ 通知退出,再 join 所有线程,最后才销毁成员)
// - ThreadPool 不可拷贝/移动(地址永不变)
workers_.emplace_back(&ThreadPool::worker_loop, this);
}
}
/**
* 析构函数:优雅关闭所有 worker。
*
* 关闭顺序:
* 1. 设置 done_ = truerelease 语义,确保 worker 可见)
* 2. notify_all() 唤醒所有正在睡眠的 worker
* 3. join 所有线程(等待它们完成当前任务并退出)
*
* 为什么要在 join 之前排空队列?
* - 如果 join 时丢弃队列中剩余任务,已提交任务的 future 将永远
* 不会 ready → 调用方 hang 住
* - worker 的 worker_loop 在退出主循环后会主动排空队列
* - join 保证了排空完成后再析构
*/
~ThreadPool() {
// release:确保 worker 的 acquire load 能看到 done_ = true
done_.store(true, std::memory_order_release);
// 唤醒所有可能在 condition_variable 上等待的 worker
cv_.notify_all();
// 等待所有 worker 退出(包括排空队列)
for (std::thread& worker : workers_) {
if (worker.joinable()) {
worker.join();
}
}
}
// ========================================================================
// 禁止拷贝/移动
// ========================================================================
ThreadPool(const ThreadPool&) = delete;
ThreadPool& operator=(const ThreadPool&) = delete;
ThreadPool(ThreadPool&&) = delete;
ThreadPool& operator=(ThreadPool&&) = delete;
// ========================================================================
// submit:提交任务,返回 future
// ========================================================================
/**
* 向线程池提交一个可调用对象及其参数。
*
* @param f 可调用对象(函数、lambda、std::bind 等)
* @param args 转发给 f 的参数
* @return std::future<R>R 是 f(args...) 的返回类型
*
* ## 实现步骤
*
* 1. 将 f(args...) 包装为 std::packaged_task<R()>(捕获调用 + 存储返回值)
* 2. 在移动 task 之前获取 futurepackaged_task::get_future() 必须
* 在 task 执行之前调用)
* 3. 类型擦除为 std::function<void()>:一个 lambda 调用 packaged_task
* 4. 入队到无锁队列中,如满则 yield 重试
* 5. notify_one() 唤醒一个可能的睡眠 worker
*
* ## 为什么需要 shared_ptr< packaged_task >
*
* std::packaged_task 是 move-only 的,但 std::function 要求可拷贝。
* shared_ptr 解决了这个矛盾:lambda 捕获 shared_ptr(拷贝),
* 内部通过指针调用 packaged_task。
*
* ## 队列满时的处理
*
* 当前实现:yield + 重试直到成功。
* 这在上游速率可控时是安全的。
* 对于生产级代码,可考虑添加超时或返回 std::optional<future>。
*/
template <typename F, typename... Args>
auto submit(F&& f, Args&&... args)
-> std::future<std::invoke_result_t<F, Args...>>
{
using result_type = std::invoke_result_t<F, Args...>;
// 将调用包装为 packaged_task。
// shared_ptr 解决 packaged_task move-only 与 function copyable 的矛盾。
auto task = std::make_shared<std::packaged_task<result_type()>>(
std::bind(std::forward<F>(f), std::forward<Args>(args)...)
);
// 获取 future。
// 必须在 task 被移动/调用之前执行。
std::future<result_type> result = task->get_future();
// 类型擦除:将 packaged_task 调用包装为 void()
std::function<void()> wrapper = [task]() {
(*task)(); // 执行 task,结果存入 packaged_task 的 shared state
};
// 入队。
// 如果队列满,yield 让出 CPU 给 workerworker 会消费任务腾出空间。
// notify_one 也可能唤醒一个正在睡眠的 worker 来帮忙消费。
while (!task_queue_.enqueue(std::move(wrapper))) {
cv_.notify_one();
std::this_thread::yield();
}
// 唤醒一个 worker(如果所有都在睡眠)。
// 不需要持有 mutex:现代实现中无锁 notify 是安全的。
cv_.notify_one();
return result;
}
private:
// ========================================================================
// worker_loopworker 线程主循环
// ========================================================================
/**
* 每个 worker 线程执行的主循环。
*
* ## 两阶段混合等待策略
*
* ### Phase 1 — 快速路径(自旋 + 无锁出队)
* 在没有获取任何锁的情况下尝试 dequeue。
* 自旋最多 100 次(经验值,在延迟和 CPU 之间折中)。
* 如果拿到任务 → 执行 → 立刻回到 Phase 1。
* 此路径在有持续负载时延迟为亚微秒级。
*
* ### Phase 2 — 慢速路径(condition_variable 等待)
* 当 Phase 1 没找到任务时:
* a) 检查 done_ 标志(如果关闭则退出主循环)
* b) 获取 mutex_
* c) 再尝试一次 dequeue(关键!防止在 Phase 1 结束和 mutex 获取之间
* 错失任务)
* d) 仍然无任务 → cv_.wait_for(1ms)
* - 超时 1ms 是安全网:即使 notify 丢失,worker 也会定期醒来检查
* - 1ms 足够长以省 CPU,足够短以保证响应
*
* ## 关闭行为
*
* 当 done_ = true 时:
* 1. 退出主循环
* 2. 排空队列中所有剩余任务(保证已提交任务不丢失)
* 3. 函数返回,线程结束
*
* ## 为什么排空在 done_ 之后?
*
* submit() 可能在 done_ 设置之前瞬间入队一个任务。
* 如果 worker 看到 done_ 就立刻退出,这个任务就会被丢弃,
* 对应的 future 永远不 ready → 调用方 hang。
* 排空阶段确保零任务丢失。
*/
void worker_loop() {
Task task;
while (true) {
// ================================================================
// Phase 1:快速路径(自旋 + 无锁出队)
// ================================================================
// 自旋 100 次是经验折中值:
// - 太小:频繁进入慢速路径,延迟增加
// - 太大:CPU 空转,浪费电力
// 在生产环境中可根据实际负载 profile 调整。
bool got_work = false;
for (int spin = 0; spin < 100; ++spin) {
if (task_queue_.dequeue(task)) {
got_work = true;
break;
}
}
if (got_work) {
task(); // 执行任务
continue; // 立刻回到 Phase 1 尝试获取更多任务
}
// ================================================================
// 检查退出条件
// ================================================================
// acquire:确保能看到 done_ = true 之前的所有 store。
if (done_.load(std::memory_order_acquire)) {
break; // 退出主循环 → 进入排空阶段
}
// ================================================================
// Phase 2:慢速路径(条件变量等待)
// ================================================================
{
std::unique_lock<std::mutex> lock(mutex_);
// 双重检查:任务可能在 Phase 1 结束到 mutex 获取之间入队。
// 如果不检查就直接 wait,可能永远收不到通知(submit 的
// notify_one 已在此检查之前发出)。
if (task_queue_.dequeue(task)) {
lock.unlock();
task();
continue;
}
// 等待通知或超时。
// 1ms 超时防止永久睡眠(notify 丢失的安全网)。
cv_.wait_for(lock, std::chrono::milliseconds(1));
}
}
// ================================================================
// 排空阶段:处理 shutdown 前提交的所有剩余任务
// ================================================================
// 此时 done_ = true,不再有新的 submit。
// 消费队列中所有剩余任务,确保零丢弃。
while (task_queue_.dequeue(task)) {
task();
}
}
// ========================================================================
// 成员变量
// ========================================================================
/// 任务类型别名:void() 类型擦除,可存储任意可调用对象。
using Task = std::function<void()>;
/// 无锁有界 MPMC 任务队列。
/// 容量 1024:在典型网络服务器场景下,足够容纳短时间内的突发任务提交,
/// 同时又足够小以提供有意义的背压(队列满即提示过载)。
LockFreeQueue<Task, 1024> task_queue_;
/// Worker 线程池。构造时预分配,整个生命周期固定不变。
std::vector<std::thread> workers_;
/// 互斥锁,仅用于 condition_variable(慢速路径)。
/// 注意:此 mutex 与无锁队列完全独立,
/// 快速路径的队列操作完全不触及此锁。
std::mutex mutex_;
/// 条件变量,用于在队列为空时让 worker 睡眠。
/// submit() 调用 notify_one() 唤醒一个 worker。
/// 析构函数调用 notify_all() 唤醒所有 worker。
std::condition_variable cv_;
/// 关闭标志。
/// 析构函数设为 true;worker 检查此标志决定是否退出。
/// atomic 保证跨线程可见性。
std::atomic<bool> done_;
};
-21
View File
@@ -1,21 +0,0 @@
#include <iostream>
#include <vector>
template <typename T>
void shifting(vector<T> &list, int k) {
for (int i = 0; (i + 1) * k < list.size(); ++ i) {
for (int j = 0; j < k; ++ j) {
std::swap(list[i * k + j], list[(i + 1) * k + j]);
}
}
int m = n % k;
for (int i = 0; i < m; ++ i) {
for (int j = 0; j < k; ++ j) {
std::swap(list[n - k], list[n - k - 1]);
}
}
}
int main() {
return 0;
}
+360
View File
@@ -0,0 +1,360 @@
/**
* @file test_pool.cpp
* @brief ThreadPool + LockFreeQueue 综合测试
*
* 测试场景:
* 1. 基本提交与返回值验证
* 2. 多生产者并发提交
* 3. 网络 I/O 场景模拟
* 4. 关闭安全性验证
*/
#include "pool.hpp"
#include <cassert>
#include <chrono>
#include <iostream>
#include <numeric>
#include <string>
#include <vector>
// ============================================================================
// 测试辅助宏
// ============================================================================
// 轻量断言,失败时打印行号
#define TEST_ASSERT(cond, msg) \
do { \
if (!(cond)) { \
std::cerr << "FAIL [" << __LINE__ << "]: " << msg << std::endl; \
return false; \
} \
} while (0)
// 运行一个测试用例并记录结果
static int g_passed = 0;
static int g_failed = 0;
#define RUN_TEST(name) \
do { \
std::cout << " " << #name << "... "; \
if (test_##name()) { \
std::cout << "PASSED" << std::endl; \
++g_passed; \
} else { \
std::cout << "FAILED" << std::endl; \
++g_failed; \
} \
} while (0)
// ============================================================================
// 测试用例
// ============================================================================
// ---------------------------------------------------------------------------
// Test 1: 基本提交 + future 返回值
// ---------------------------------------------------------------------------
static bool test_basic_submit() {
ThreadPool pool(4);
// 提交一个返回 int 的任务
auto f1 = pool.submit([] { return 42; });
TEST_ASSERT(f1.get() == 42, "simple int return");
// 提交带参数的任务
auto f2 = pool.submit([](int a, int b) { return a + b; }, 10, 32);
TEST_ASSERT(f2.get() == 42, "parameterized callable");
// 提交返回 string 的任务
auto f3 = pool.submit([](const std::string& s) { return s + " world"; },
"hello");
TEST_ASSERT(f3.get() == "hello world", "string return");
// 提交 void 任务
bool called = false;
auto f4 = pool.submit([&called] { called = true; });
f4.get(); // wait for completion
TEST_ASSERT(called, "void task should set flag");
return true;
}
// ---------------------------------------------------------------------------
// Test 2: 多生产者并发提交
// ---------------------------------------------------------------------------
static bool test_multi_producer() {
constexpr int num_producers = 8;
constexpr int tasks_per_producer = 500;
ThreadPool pool(4);
// 所有 producer 向同一个计数器累加
std::atomic<long long> counter{0};
// 启动多个 producer 线程,每个提交若干任务
std::vector<std::thread> producers;
producers.reserve(num_producers);
for (int p = 0; p < num_producers; ++p) {
producers.emplace_back([&pool, &counter, tasks_per_producer, p] {
std::vector<std::future<void>> futures;
for (int i = 0; i < tasks_per_producer; ++i) {
auto f = pool.submit([&counter] {
// 原子加 1
counter.fetch_add(1, std::memory_order_relaxed);
});
futures.push_back(std::move(f));
}
// 等待此 producer 的所有任务完成
for (auto& f : futures) {
f.get();
}
});
}
// 加入所有 producer
for (auto& t : producers) {
t.join();
}
// 验证总计数
long long expected = static_cast<long long>(num_producers) * tasks_per_producer;
TEST_ASSERT(counter.load() == expected,
"counter should match total tasks: " << counter.load()
<< " vs " << expected);
return true;
}
// ---------------------------------------------------------------------------
// Test 3: 网络 I/O 场景模拟
// ---------------------------------------------------------------------------
static bool test_network_scenario() {
/*
* 模拟典型的网络服务器任务模式:
* - 多个连接并发到达
* - 每个连接经过:接收 → 处理 → 响应 流水线
* - 使用 future 链接后续处理步骤
*/
ThreadPool pool(std::thread::hardware_concurrency());
constexpr int num_connections = 200;
std::atomic<int> received_count{0};
std::atomic<int> processed_count{0};
std::atomic<int> responded_count{0};
std::vector<std::future<void>> final_futures;
final_futures.reserve(num_connections);
for (int conn_id = 0; conn_id < num_connections; ++conn_id) {
// 模拟处理一个连接:接收 → 处理 → 响应
auto f = pool.submit([&received_count, &processed_count, &responded_count, conn_id] {
// Phase 1: 模拟接收数据
received_count.fetch_add(1, std::memory_order_relaxed);
// Phase 2: 模拟业务处理
// 计算一些东西来表示处理
volatile int result = 0;
for (int i = 0; i < 100; ++i) {
result += (conn_id ^ i) & 0xFF;
}
(void)result;
processed_count.fetch_add(1, std::memory_order_relaxed);
// Phase 3: 模拟发送响应
responded_count.fetch_add(1, std::memory_order_relaxed);
});
final_futures.push_back(std::move(f));
}
// 等待所有连接处理完毕
for (auto& f : final_futures) {
f.get();
}
TEST_ASSERT(received_count.load() == num_connections,
"all connections received");
TEST_ASSERT(processed_count.load() == num_connections,
"all connections processed");
TEST_ASSERT(responded_count.load() == num_connections,
"all connections responded");
return true;
}
// ---------------------------------------------------------------------------
// Test 4: 关闭安全性 — 析构时排空所有已提交任务
// ---------------------------------------------------------------------------
static bool test_shutdown_drain() {
constexpr int num_tasks = 500;
std::atomic<int> completed{0};
{
ThreadPool pool(4);
// 提交大量任务(不获取 future,依赖析构排空)
for (int i = 0; i < num_tasks; ++i) {
// 不使用返回值,仅依赖析构保证执行
auto f = pool.submit([&completed] {
completed.fetch_add(1, std::memory_order_relaxed);
});
(void)f; // 忽略 future,依赖析构
}
// pool 在此作用域结束析构,排空所有任务
}
// 析构后全部任务应已完成
TEST_ASSERT(completed.load() == num_tasks,
"all tasks drained on shutdown: " << completed.load()
<< " vs " << num_tasks);
return true;
}
// ---------------------------------------------------------------------------
// Test 5: LockFreeQueue 基本功能测试
// ---------------------------------------------------------------------------
static bool test_lockfree_queue() {
LockFreeQueue<int, 8> q;
// 空队列状态
TEST_ASSERT(q.empty(), "queue should be empty initially");
TEST_ASSERT(!q.full(), "queue should not be full initially");
TEST_ASSERT(q.size() == 0, "size should be 0 initially");
// 入队直到满
for (int i = 0; i < 8; ++i) {
TEST_ASSERT(q.enqueue(std::move(i)), "enqueue should succeed");
}
TEST_ASSERT(q.full(), "queue should be full after 8 enqueues");
// 入队失败(队列满)
int extra = 99;
TEST_ASSERT(!q.enqueue(std::move(extra)), "enqueue should fail when full");
// 出队所有元素
for (int i = 0; i < 8; ++i) {
int val = -1;
TEST_ASSERT(q.dequeue(val), "dequeue should succeed");
TEST_ASSERT(val == i, "dequeued value should match: " << val << " vs " << i);
}
TEST_ASSERT(q.empty(), "queue should be empty after 8 dequeues");
// 出队失败(队列空)
int val = -1;
TEST_ASSERT(!q.dequeue(val), "dequeue should fail when empty");
return true;
}
// ---------------------------------------------------------------------------
// Test 6: LockFreeQueue 多线程压力测试
// ---------------------------------------------------------------------------
static bool test_lockfree_queue_threaded() {
constexpr size_t Q_CAP = 256;
LockFreeQueue<int, Q_CAP> q;
constexpr int num_producers = 4;
constexpr int num_consumers = 4;
constexpr int items_per_producer = 10000;
constexpr long long total_items = static_cast<long long>(num_producers) * items_per_producer;
// 每个 producer 入队的起始值(按 producer 编号偏移,避免重复)
std::atomic<long long> enqueued_sum{0};
std::atomic<long long> dequeued_sum{0};
std::atomic<long long> enqueued_count{0};
std::atomic<long long> dequeued_count{0};
std::atomic<bool> producers_done{false};
// 启动 producers
std::vector<std::thread> producers;
for (int p = 0; p < num_producers; ++p) {
producers.emplace_back([&q, &enqueued_sum, &enqueued_count, items_per_producer, p] {
for (int i = 0; i < items_per_producer; ++i) {
int val = p * items_per_producer + i;
while (!q.enqueue(std::move(val))) {
std::this_thread::yield();
}
enqueued_sum.fetch_add(val, std::memory_order_relaxed);
enqueued_count.fetch_add(1, std::memory_order_relaxed);
}
});
}
// 启动 consumers
std::vector<std::thread> consumers;
for (int c = 0; c < num_consumers; ++c) {
consumers.emplace_back([&q, &dequeued_sum, &dequeued_count, &producers_done] {
while (true) {
int val = -1;
if (q.dequeue(val)) {
dequeued_sum.fetch_add(val, std::memory_order_relaxed);
dequeued_count.fetch_add(1, std::memory_order_relaxed);
} else if (producers_done.load(std::memory_order_acquire)) {
// 生产者完成且队列空 → 退出
// 再尝试一次(可能在 producers_done 设置前瞬间入队)
if (q.dequeue(val)) {
dequeued_sum.fetch_add(val, std::memory_order_relaxed);
dequeued_count.fetch_add(1, std::memory_order_relaxed);
continue;
}
break;
}
}
});
}
// 等待所有 producer 完成
for (auto& t : producers) {
t.join();
}
producers_done.store(true, std::memory_order_release);
// 等待所有 consumer 完成
for (auto& t : consumers) {
t.join();
}
// 验证
TEST_ASSERT(enqueued_count.load() == total_items,
"all items enqueued");
TEST_ASSERT(dequeued_count.load() == total_items,
"all items dequeued: " << dequeued_count.load() << " vs " << total_items);
TEST_ASSERT(enqueued_sum.load() == dequeued_sum.load(),
"sums match: " << enqueued_sum.load() << " vs " << dequeued_sum.load());
TEST_ASSERT(q.empty(), "queue should be empty at end");
return true;
}
// ============================================================================
// main
// ============================================================================
int main() {
std::cout << "=== ThreadPool & LockFreeQueue Tests ===" << std::endl;
std::cout << "Hardware concurrency: "
<< std::thread::hardware_concurrency() << std::endl;
std::cout << std::endl;
std::cout << "[LockFreeQueue]" << std::endl;
RUN_TEST(lockfree_queue);
RUN_TEST(lockfree_queue_threaded);
std::cout << std::endl;
std::cout << "[ThreadPool]" << std::endl;
RUN_TEST(basic_submit);
RUN_TEST(multi_producer);
RUN_TEST(network_scenario);
RUN_TEST(shutdown_drain);
std::cout << std::endl;
std::cout << "=== Results: " << g_passed << " passed, "
<< g_failed << " failed ===" << std::endl;
return (g_failed == 0) ? 0 : 1;
}