feat(repair): 断点续跑 progress 文件管理
load_progress / save_progress(asyncio.Lock + os.replace 原子写入) / should_skip_video。支持并发安全的读改写和 --reaggregate-all 兜底。 Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,71 @@
|
||||
"""修复管线断点续跑 progress 管理测试。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def test_load_progress_missing_file(tmp_path):
|
||||
"""progress 文件不存在时返回空集合。"""
|
||||
from tools.repair_trees import load_progress
|
||||
result = load_progress(tmp_path / "nonexistent.json")
|
||||
assert result == set()
|
||||
|
||||
|
||||
def test_load_progress_valid_file(tmp_path):
|
||||
"""正常读取已有 progress 文件。"""
|
||||
from tools.repair_trees import load_progress
|
||||
path = tmp_path / "progress.json"
|
||||
path.write_text(json.dumps({"finished_video_ids": ["vid_a", "vid_b"]}))
|
||||
result = load_progress(path)
|
||||
assert result == {"vid_a", "vid_b"}
|
||||
|
||||
|
||||
def test_load_progress_corrupted_file(tmp_path):
|
||||
"""损坏的 JSON 文件返回空集合(不抛异常)。"""
|
||||
from tools.repair_trees import load_progress
|
||||
path = tmp_path / "progress.json"
|
||||
path.write_text("{invalid json")
|
||||
result = load_progress(path)
|
||||
assert result == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_progress_atomic(tmp_path):
|
||||
"""save_progress 原子写入,并发调用不丢失更新。"""
|
||||
from tools.repair_trees import save_progress
|
||||
path = tmp_path / "progress.json"
|
||||
lock = asyncio.Lock()
|
||||
await save_progress(path, lock, "vid_a")
|
||||
await save_progress(path, lock, "vid_b")
|
||||
data = json.loads(path.read_text())
|
||||
assert set(data["finished_video_ids"]) == {"vid_a", "vid_b"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_progress_concurrent(tmp_path):
|
||||
"""16 路并发 save_progress 不丢失更新。"""
|
||||
from tools.repair_trees import save_progress
|
||||
path = tmp_path / "progress.json"
|
||||
lock = asyncio.Lock()
|
||||
tasks = [save_progress(path, lock, f"vid_{i}") for i in range(16)]
|
||||
await asyncio.gather(*tasks)
|
||||
data = json.loads(path.read_text())
|
||||
assert len(data["finished_video_ids"]) == 16
|
||||
|
||||
|
||||
def test_should_skip_finished():
|
||||
"""已在 finished 集合中的视频应跳过。"""
|
||||
from tools.repair_trees import should_skip_video
|
||||
finished = {"vid_a", "vid_b"}
|
||||
assert should_skip_video("vid_a", finished, reaggregate_all=False) is True
|
||||
assert should_skip_video("vid_c", finished, reaggregate_all=False) is False
|
||||
|
||||
|
||||
def test_should_skip_reaggregate_all_forces_rerun():
|
||||
"""--reaggregate-all 标志强制不跳过。"""
|
||||
from tools.repair_trees import should_skip_video
|
||||
finished = {"vid_a"}
|
||||
assert should_skip_video("vid_a", finished, reaggregate_all=True) is False
|
||||
Reference in New Issue
Block a user