"""nodes/adaptive_pool.py 的模块级测试(数据 → 测试过程 → 验证结果)。 被测模块:`nodes/adaptive_pool.py`(自适应并发额度控制),被 subtitle-ocr 与 llm-filter 复用,也可独立使用,因此拥有独立模块目录。 测试只调用真实并发池,`clock` 注入假时钟以实现确定性的窗口行为;worker 用 真实可执行函数(无 I/O 依赖),不重写被测逻辑。 """ from __future__ import annotations import threading import time from nodes.adaptive_pool import AdaptiveThreadPool, decide class FakeClock: """可手动拨动的假时钟(时间属于允许在 I/O 边界注入的依赖)。""" def __init__(self, now: float = 0.0) -> None: self.now = now def __call__(self) -> float: return self.now def advance(self, seconds: float) -> None: self.now += seconds # --------------------------------------------------------------------------- # 决策函数 # --------------------------------------------------------------------------- def test_decide_increases_when_response_fast() -> None: """平均响应低于快阈值且未达上限:额度 +1。""" # 数据:当前 1、均值 0.1、上限 16。 # 测试过程与验证结果 assert decide(1, 0.1, 1, 16, 0.3, 1.0) == 2 def test_decide_decreases_when_response_slow() -> None: """平均响应高于慢阈值且高于下限:额度 -1。""" # 数据:当前 3、均值 2.0。 # 测试过程与验证结果 assert decide(3, 2.0, 1, 16, 0.3, 1.0) == 2 def test_decide_keeps_when_response_between_thresholds() -> None: """响应介于两阈值之间时额度不变。""" # 数据:当前 2、均值 0.5。 # 测试过程与验证结果 assert decide(2, 0.5, 1, 16, 0.3, 1.0) == 2 def test_decide_respects_bounds() -> None: """已达上限不再增、已达下限不再减。""" # 数据:上限 16 且很快;下限 1 且很慢。 # 测试过程与验证结果 assert decide(16, 0.1, 1, 16, 0.3, 1.0) == 16 assert decide(1, 2.0, 1, 16, 0.3, 1.0) == 1 # --------------------------------------------------------------------------- # map 基本行为 # --------------------------------------------------------------------------- def test_map_returns_results_in_input_order() -> None: """并发执行但结果严格按输入顺序返回,保证字幕时间轴不被并发打乱。""" # 数据:让后面的任务更早开始的输入;worker 是确定性的纯函数。 items = [3, 2, 1] pool = AdaptiveThreadPool(worker=lambda value: value * 100) # 测试过程 results = pool.map(items) # 验证结果:顺序与输入一致。 assert results == [300, 200, 100] def test_map_empty_input_returns_empty_list() -> None: """空输入返回空结果,不启动任务也不报错。""" # 数据:空列表。 pool = AdaptiveThreadPool(worker=lambda item: item) # 测试过程与验证结果 assert pool.map([]) == [] def test_map_isolates_worker_exception_as_result() -> None: """worker 抛出的异常作为该位置的结果返回,不影响其他任务。""" # 数据:第二个元素触发异常。 def worker(item: int) -> int: if item == 2: raise ValueError("boom") return item pool = AdaptiveThreadPool(worker=worker) # 测试过程 results = pool.map([1, 2, 3]) # 验证结果:异常对象保留在对应位置,其余结果正常。 assert results[0] == 1 assert isinstance(results[1], ValueError) assert results[2] == 3 def test_progress_callback_reports_completed_and_workers() -> None: """进度回调按完成数递增上报,并带上当前额度。""" # 数据:记录全部回调的列表。 calls: list[tuple[int, int, int]] = [] pool = AdaptiveThreadPool( worker=lambda item: item, on_progress=lambda done, total, rate, avg, workers: calls.append((done, total, workers)), ) # 测试过程 pool.map([1, 2, 3]) # 验证结果:完成数依次为 1、2、3,总数恒为 3,额度不低于下限。 assert [c[0] for c in calls] == [1, 2, 3] assert all(c[1] == 3 for c in calls) assert all(c[2] >= 1 for c in calls) def test_cancel_suppresses_progress_but_supplies_all_results() -> None: """cancel 只抑制进度回调;每个输入仍返回结果(供调用方决定重试)。""" # 数据:worker 在第一个任务里触发 cancel。 pool: AdaptiveThreadPool def worker(item: int) -> int: if item == 0: pool.cancel() return item pool = AdaptiveThreadPool(worker=worker, on_progress=lambda *a: calls.append(a)) calls: list[tuple] = [] # 测试过程 results = pool.map([0, 1, 2]) # 验证结果:结果完整,进度回调被抑制。 assert results == [0, 1, 2] assert calls == [] def test_cancel_is_reset_for_next_map() -> None: """同一实例再次 map 时清除 cancel 标记(暂停后继续、失败重试的调用约定)。""" # 数据:第一批触发 cancel,第二批正常。 pool = AdaptiveThreadPool(worker=lambda item: item) pool.cancel() calls: list[tuple] = [] pool._on_progress = lambda *a: calls.append(a) # 直接复用真实回调槽位 # 测试过程 pool.map([1]) # 验证结果:新批次重新上报进度。 assert len(calls) == 1 # --------------------------------------------------------------------------- # 弹性扩缩容(假时钟驱动确定性窗口) # --------------------------------------------------------------------------- def _timed_pool(clock: FakeClock, worker, **kwargs) -> AdaptiveThreadPool: """构造使用假时钟的真实线程池。""" return AdaptiveThreadPool(worker=worker, clock=clock, window_seconds=1.0, **kwargs) def test_pool_grows_target_when_responses_fast() -> None: """窗口内平均响应快时应扩大目标额度(服务端空闲就加大并发)。""" # 数据:真实时钟 + 极短窗口;worker 只做微秒级工作,平均耗时远低于快阈值。 observed: list[int] = [] def worker(item: int) -> int: time.sleep(0.002) return item pool = AdaptiveThreadPool( worker=worker, min_workers=1, max_workers=8, window_seconds=0.01, fast_threshold=0.3, slow_threshold=1.0, ) pool._on_progress = lambda done, total, rate, avg, workers: observed.append(workers) # 测试过程:任务足够多,保证跨越多个窗口触发扩容判断。 pool.map(list(range(40))) # 验证结果:额度单调不减且最终大于起始值。 assert observed == sorted(observed) assert observed[-1] > observed[0] def test_pool_shrinks_target_when_responses_slow() -> None: """窗口内平均响应慢时应收缩目标额度(避免压垮本地服务)。""" # 数据:真实时钟 + 极短窗口;worker 耗时远超慢阈值。 observed: list[int] = [] def worker(item: int) -> int: time.sleep(0.02) return item pool = AdaptiveThreadPool( worker=worker, min_workers=1, max_workers=8, window_seconds=0.05, fast_threshold=0.001, slow_threshold=0.005, ) pool._on_progress = lambda done, total, rate, avg, workers: observed.append(workers) # 测试过程 pool.map(list(range(12))) # 验证结果:出现回调,且额度始终不低于下限(不会缩到 0)。 assert observed assert min(observed) >= 1 def test_report_failure_lowers_effective_max_and_never_expands() -> None: """report_failure 只收紧上限;上限从 20 降到 19 时不会把当前 1 并发扩成 19。""" # 数据:上限 20、当前额度 1。 pool = AdaptiveThreadPool(worker=lambda item: item, min_workers=1, max_workers=20) # 测试过程:报告一次限流失败。 pool.report_failure() # 验证结果:有效上限降 1,当前目标仍为下限。 assert pool._effective_max_workers == 19 assert pool._target_workers <= 1 def test_failure_at_single_worker_never_expands_quota() -> None: """单并发下报告失败,绝不能因上限下降而把额度扩大。""" # 数据:上限 20、最小 1,先启动一次 map 使额度为 1。 pool = AdaptiveThreadPool(worker=lambda item: item, min_workers=1, max_workers=20) pool.map([1]) # 测试过程 pool.report_failure() # 验证结果 assert pool._target_workers == 1 def test_effective_max_recovers_one_step_per_clean_window() -> None: """干净窗口每次只恢复 1 个上限,避免限流恢复期再次打满。""" # 数据:先把有效上限压到 2。 clock = FakeClock() pool = _timed_pool(clock, lambda item: item, min_workers=1, max_workers=5) pool.report_failure() pool.report_failure() pool.report_failure() assert pool._effective_max_workers == 2 # 测试过程:第一个 tick 消耗“有失败”的窗口(不恢复),第二个干净窗口恢复 1。 clock.advance(2.0) pool._tick(0.1) assert pool._effective_max_workers == 2 clock.advance(2.0) pool._tick(0.1) # 验证结果:干净窗口只恢复 1。 assert pool._effective_max_workers == 3 def test_error_window_does_not_recover_limit() -> None: """窗口内有失败时不恢复上限。""" # 数据:上限 5,压到 3 后在同一窗口内报告失败。 clock = FakeClock() pool = _timed_pool(clock, lambda item: item, min_workers=1, max_workers=5) pool.report_failure() pool.report_failure() pool.report_failure() before = pool._effective_max_workers # 测试过程:窗口内先失败再触发 tick。 pool.report_failure() clock.advance(2.0) pool._tick(0.1) # 验证结果:上限未恢复。 assert pool._effective_max_workers <= before def test_reduced_limit_is_kept_across_retry_map() -> None: """限流后的有效上限跨 map 保留(失败条目重试时继续遵守更严配额)。""" # 数据:先压低上限。 pool = AdaptiveThreadPool(worker=lambda item: item, min_workers=1, max_workers=6) pool.map([1]) pool.report_failure() reduced = pool._effective_max_workers # 测试过程:执行第二轮 map(重试场景)。 pool.map([1, 2]) # 验证结果:上限仍是被压低的值。 assert pool._effective_max_workers == reduced def test_concurrent_map_on_same_pool_is_rejected() -> None: """同一实例不允许并行 map,避免额度与统计互相干扰。""" # 数据:第一个 map 阻塞在 worker 上。 started = threading.Event() release = threading.Event() def worker(item: int) -> int: started.set() release.wait(timeout=5) return item pool = AdaptiveThreadPool(worker=worker, min_workers=1, max_workers=2) errors: list[Exception] = [] def run_first() -> None: pool.map([1]) thread = threading.Thread(target=run_first) thread.start() assert started.wait(timeout=5) # 测试过程:在第一个 map 未结束时再次调用 map。 try: pool.map([2]) except Exception as exc: # noqa: BLE001 - 断言真实抛出的类型 errors.append(exc) finally: release.set() thread.join(timeout=5) # 验证结果:抛出 RuntimeError。 assert len(errors) == 1 assert isinstance(errors[0], RuntimeError) def test_max_concurrency_never_exceeds_target_limit() -> None: """实际在途数量不超过 max_workers 上限(并发有界)。""" # 数据:上限 3,10 个快速任务。 pool = AdaptiveThreadPool(worker=lambda item: item, min_workers=1, max_workers=3) # 测试过程 pool.map(list(range(10))) # 验证结果:观测到的最大在途数不超过上限。 assert 1 <= pool.max_concurrency <= 3