diff --git a/CMakeLists.txt b/CMakeLists.txt index 66e6b2c..f4c0a61 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -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 . ) diff --git a/pool.cpp b/pool.cpp index 918a31b..b14d807 100644 --- a/pool.cpp +++ b/pool.cpp @@ -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 queue; - - queue.push(1); - auto load = queue.pop(); -} diff --git a/pool.hpp b/pool.hpp index cace14d..b9a5672 100644 --- a/pool.hpp +++ b/pool.hpp @@ -1,44 +1,677 @@ -#include +#pragma once + +/** + * @file pool.hpp + * @brief 无锁有界 MPMC 队列 + 固定大小线程池,适用于多线程网络通信场景。 + * + * 本文件提供两个组件: + * 1. LockFreeQueue — 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 -#include #include +#include #include - -#include - +#include +#include +#include +#include +#include +#include #include -template -struct QNode { - std::unique_ptr payload; - volatile std::atomic next; +#include - QNode() = default; +// ============================================================================ +// LockFreeQueue — Vyukov 有界 MPMC 无锁队列 +// ============================================================================ - explicit QNode(T&& payload): - payload(std::forward(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 +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 绕过整个环回到 + * 同一 cell,sequence 的值也不同,CAS 永远不会在错误回合成功。 + * 这就是「天然免疫 ABA」的原因。 + */ + struct Cell { + std::atomic sequence; + T data; + + Cell() noexcept : sequence(0) {} + }; + + // ======================================================================== + // 环形缓冲区 + // ======================================================================== + 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 head_; + alignas(64) std::atomic tail_; -template -class Queue { -private: - volatile std::atomic*> head; - volatile std::atomic*> tail; public: - void push(T&& payload) { - QNode* node = new QNode(std::forward(payload)); - node->next = node; + // ======================================================================== + // 构造 + // ======================================================================== - QNode* _tail = tail; - QNode* _next = _tail->next; - - do { - node->next = _next; - } while (!_tail->next.compare_exchange_weak(_next, node)); - - tail.compare_exchange_strong(_tail, node); + /** + * 初始化所有 cell:sequence[i] = i。 + * + * 为什么 sequence[i] = i? + * + * 构造时,head = 0, tail = 0。 + * cell[0].sequence = 0 → 对于 pos=0,seq == pos(空闲可写) + * cell[1].sequence = 1 → 对于 pos=1,seq == 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 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(seq) - static_cast(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(seq) - static_cast(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_ = true(release 语义,确保 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 是 f(args...) 的返回类型 + * + * ## 实现步骤 + * + * 1. 将 f(args...) 包装为 std::packaged_task(捕获调用 + 存储返回值) + * 2. 在移动 task 之前获取 future(packaged_task::get_future() 必须 + * 在 task 执行之前调用) + * 3. 类型擦除为 std::function:一个 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。 + */ + template + auto submit(F&& f, Args&&... args) + -> std::future> + { + using result_type = std::invoke_result_t; + + // 将调用包装为 packaged_task。 + // shared_ptr 解决 packaged_task move-only 与 function copyable 的矛盾。 + auto task = std::make_shared>( + std::bind(std::forward(f), std::forward(args)...) + ); + + // 获取 future。 + // 必须在 task 被移动/调用之前执行。 + std::future result = task->get_future(); + + // 类型擦除:将 packaged_task 调用包装为 void() + std::function wrapper = [task]() { + (*task)(); // 执行 task,结果存入 packaged_task 的 shared state + }; + + // 入队。 + // 如果队列满,yield 让出 CPU 给 worker,worker 会消费任务腾出空间。 + // 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_loop:worker 线程主循环 + // ======================================================================== + + /** + * 每个 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 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; + + /// 无锁有界 MPMC 任务队列。 + /// 容量 1024:在典型网络服务器场景下,足够容纳短时间内的突发任务提交, + /// 同时又足够小以提供有意义的背压(队列满即提示过载)。 + LockFreeQueue task_queue_; + + /// Worker 线程池。构造时预分配,整个生命周期固定不变。 + std::vector workers_; + + /// 互斥锁,仅用于 condition_variable(慢速路径)。 + /// 注意:此 mutex 与无锁队列完全独立, + /// 快速路径的队列操作完全不触及此锁。 + std::mutex mutex_; + + /// 条件变量,用于在队列为空时让 worker 睡眠。 + /// submit() 调用 notify_one() 唤醒一个 worker。 + /// 析构函数调用 notify_all() 唤醒所有 worker。 + std::condition_variable cv_; + + /// 关闭标志。 + /// 析构函数设为 true;worker 检查此标志决定是否退出。 + /// atomic 保证跨线程可见性。 + std::atomic done_; +}; diff --git a/shift.cpp b/shift.cpp deleted file mode 100644 index b232b41..0000000 --- a/shift.cpp +++ /dev/null @@ -1,21 +0,0 @@ -#include -#include - -template -void shifting(vector &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; -} diff --git a/test_pool.cpp b/test_pool.cpp new file mode 100644 index 0000000..0328055 --- /dev/null +++ b/test_pool.cpp @@ -0,0 +1,360 @@ +/** + * @file test_pool.cpp + * @brief ThreadPool + LockFreeQueue 综合测试 + * + * 测试场景: + * 1. 基本提交与返回值验证 + * 2. 多生产者并发提交 + * 3. 网络 I/O 场景模拟 + * 4. 关闭安全性验证 + */ + +#include "pool.hpp" + +#include +#include +#include +#include +#include +#include + +// ============================================================================ +// 测试辅助宏 +// ============================================================================ + +// 轻量断言,失败时打印行号 +#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 counter{0}; + + // 启动多个 producer 线程,每个提交若干任务 + std::vector producers; + producers.reserve(num_producers); + + for (int p = 0; p < num_producers; ++p) { + producers.emplace_back([&pool, &counter, tasks_per_producer, p] { + std::vector> 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(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 received_count{0}; + std::atomic processed_count{0}; + std::atomic responded_count{0}; + + std::vector> 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 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 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 q; + + constexpr int num_producers = 4; + constexpr int num_consumers = 4; + constexpr int items_per_producer = 10000; + constexpr long long total_items = static_cast(num_producers) * items_per_producer; + + // 每个 producer 入队的起始值(按 producer 编号偏移,避免重复) + std::atomic enqueued_sum{0}; + std::atomic dequeued_sum{0}; + std::atomic enqueued_count{0}; + std::atomic dequeued_count{0}; + std::atomic producers_done{false}; + + // 启动 producers + std::vector 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 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; +}