diff --git a/tests/unit/test_repair_progress.py b/tests/unit/test_repair_progress.py new file mode 100644 index 0000000..d3a5532 --- /dev/null +++ b/tests/unit/test_repair_progress.py @@ -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 diff --git a/tools/repair_trees.py b/tools/repair_trees.py new file mode 100644 index 0000000..13d287b --- /dev/null +++ b/tools/repair_trees.py @@ -0,0 +1,430 @@ +#!/usr/bin/env python3 +"""树修复管线:检测 + VLM 重生成 + 校验 + Q&A 反向补全。 + +对 store/videos/ 下所有已迁移的树执行完整修复流程: + 1. detect_issues() — 扫描空字段/缺失帧 + 2. repair_tree() — VLM 重新描述 + 底向上级联(如有问题节点) + 3. verify_tree() — 交叉校验删除幻觉 + 4. supplement_tree() — Q&A 反向补全注入缺失事实 + 5. save_json() — 覆盖保存 + +用法: + conda activate Video-Tree-TRM + python tools/repair_trees.py [--videos-dir store/videos] [--concurrency 4] [--dry-run] + +app/core/adapters 不 import 此脚本。 +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import os +import sys +import time +from pathlib import Path + +# 确保项目根目录在 sys.path 中 +PROJECT_ROOT = Path(__file__).resolve().parent.parent +sys.path.insert(0, str(PROJECT_ROOT)) + +from dotenv import load_dotenv +from loguru import logger + +load_dotenv(PROJECT_ROOT / ".env") + +from app.tree.index import TreeIndex +from app.tree.repair.detector import detect_issues +from app.tree.repair.regenerator import repair_tree +from app.tree.repair.supplement import supplement_tree +from app.tree.subtitle import SRTEntry, parse_srt +from app.tree.verify import verify_tree + +# --------------------------------------------------------------------------- +# 日志配置:不缓存,立即输出 +# --------------------------------------------------------------------------- + +logger.remove() +logger.add( + sys.stderr, + format="{time:HH:mm:ss} | {level:<7} | {message}", + level="DEBUG", + colorize=True, +) +logger.add( + PROJECT_ROOT / "logs" / "repair_trees.log", + format="{time:YYYY-MM-DD HH:mm:ss} | {level:<7} | {message}", + level="DEBUG", + rotation="50 MB", +) + +# --------------------------------------------------------------------------- +# 断点续跑 — progress 文件管理 +# --------------------------------------------------------------------------- + +PROGRESS_FILE = "repair_progress.json" + + +def load_progress(path: Path) -> set[str]: + """读取 progress 文件,返回已完成视频 ID 集合。 + + 参数: + path: progress JSON 文件路径。 + + 返回: + 已完成视频 ID 集合。文件不存在或损坏时返回空集。 + """ + if not path.exists(): + return set() + try: + data = json.loads(path.read_text(encoding="utf-8")) + return set(data.get("finished_video_ids", [])) + except (json.JSONDecodeError, KeyError, TypeError): + logger.warning("progress 文件损坏,忽略: {}", path) + return set() + + +async def save_progress(path: Path, lock: asyncio.Lock, vid: str) -> None: + """原子追加一个视频 ID 到 progress 文件。 + + 参数: + path: progress JSON 文件路径。 + lock: asyncio.Lock,防并发读改写丢更新。 + vid: 要追加的视频 ID。 + """ + async with lock: + finished = load_progress(path) + finished.add(vid) + tmp = path.with_suffix(".tmp") + tmp.write_text( + json.dumps({"finished_video_ids": sorted(finished)}, ensure_ascii=False, indent=2), + encoding="utf-8", + ) + os.replace(str(tmp), str(path)) + + +def should_skip_video(vid: str, finished: set[str], *, reaggregate_all: bool) -> bool: + """判断是否跳过该视频。 + + 参数: + vid: 视频 ID。 + finished: progress 中已完成的视频 ID 集合。 + reaggregate_all: --reaggregate-all 标志。 + + 返回: + True 表示跳过。 + """ + if reaggregate_all: + return False + return vid in finished + + +# --------------------------------------------------------------------------- +# LLM/VLM 客户端构建 +# --------------------------------------------------------------------------- + + +def _build_clients(): + """构建 GovernedLLMClient(LLM + VLM)。 + + 返回: + (llm_client, vlm_client) 元组。 + """ + from adapters.breaker import CircuitBreaker + from adapters.llm import GovernedLLMClient + from adapters.telemetry import SQLiteTelemetryRecorder + from adapters.vlm import GovernedVLMClient + + # 遥测记录器(GovernedLLMClient 要求非 None) + (PROJECT_ROOT / "logs").mkdir(exist_ok=True) + telemetry = SQLiteTelemetryRecorder(str(PROJECT_ROOT / "logs" / "repair_telemetry.db")) + + breaker_threshold = int(os.getenv("LLM_CIRCUIT_BREAKER_THRESHOLD", "5")) + breaker_cooldown = int(os.getenv("LLM_CIRCUIT_BREAKER_COOLDOWN", "60")) + timeout_s = float(os.getenv("LLM_TIMEOUT", "120")) + max_retries = int(os.getenv("LLM_MAX_RETRIES", "3")) + base_delay = float(os.getenv("LLM_RETRY_BASE_DELAY", "2.0")) + max_delay = float(os.getenv("LLM_RETRY_MAX_DELAY", "30.0")) + ttft = float(os.getenv("LLM_TTFT_TIMEOUT", "30")) + inter_token = float(os.getenv("LLM_INTER_TOKEN_TIMEOUT", "15")) + + # LLM 客户端(用于 supplement 和 L2/L1 重生成) + llm = GovernedLLMClient( + model=os.environ["SEARCH_LLM_MODEL"], + base_url=os.environ["SEARCH_LLM_BASE_URL"], + api_key=os.environ["SEARCH_LLM_API_KEY"], + provider="deepseek", + thinking=False, + breaker=CircuitBreaker(fail_threshold=breaker_threshold, cooldown_s=breaker_cooldown), + cache=None, + telemetry=telemetry, + timeout_s=timeout_s, + ttft_timeout_s=ttft, + inter_token_timeout_s=inter_token, + max_retries=max_retries, + retry_base_delay_s=base_delay, + retry_max_delay_s=max_delay, + ) + + # VLM 客户端(用于 L3 帧重新描述) + vlm_base = GovernedLLMClient( + model=os.environ["VL_LLM_MODEL"], + base_url=os.environ["VL_LLM_BASE_URL"], + api_key=os.environ["VL_LLM_API_KEY"], + provider="qwen", + thinking=False, + breaker=CircuitBreaker(fail_threshold=breaker_threshold, cooldown_s=breaker_cooldown), + cache=None, + telemetry=telemetry, + timeout_s=timeout_s, + ttft_timeout_s=ttft, + inter_token_timeout_s=inter_token, + max_retries=max_retries, + retry_base_delay_s=base_delay, + retry_max_delay_s=max_delay, + ) + vlm = GovernedVLMClient(vlm_base) + + return llm, vlm + + +# --------------------------------------------------------------------------- +# 单视频修复 +# --------------------------------------------------------------------------- + + +async def _repair_one_video( + vid: str, + tree_path: Path, + frames_dir: Path, + srt_dir: Path, + questions_dir: Path, + llm, + vlm, + *, + dry_run: bool = False, +) -> dict: + """修复单个视频的树。 + + 参数: + vid: 视频 ID。 + tree_path: tree.json 路径。 + frames_dir: 帧文件目录。 + srt_dir: SRT 字幕目录。 + questions_dir: 问题 JSON 目录。 + llm: LLMProvider 实例。 + vlm: VLMProvider 实例。 + dry_run: 仅检测不修复。 + + 返回: + 统计 dict。 + """ + stats = { + "vid": vid, + "issues_found": 0, + "l3_repaired": 0, + "l2_regenerated": 0, + "l1_regenerated": 0, + "verify_removed": 0, + "facts_injected": 0, + "error": None, + } + + try: + # 加载树 + index = TreeIndex.load_json(str(tree_path)) + + # Step 1: 检测问题 + issues = detect_issues(index, frames_dir=frames_dir) + stats["issues_found"] = len(issues) + + if issues: + logger.info("[{}] 发现 {} 个问题", vid, len(issues)) + for issue in issues[:5]: + logger.debug(" {} [L{}] {}", issue.node_id, issue.level, issue.details) + if len(issues) > 5: + logger.debug(" ... 还有 {} 个", len(issues) - 5) + + if dry_run: + return stats + + # Step 2: VLM 修复(如有 empty_field 问题) + empty_issues = [i for i in issues if i.issue_type == "empty_field"] + if empty_issues: + srt_entries = None + srt_path = srt_dir / f"{vid}.srt" + if srt_path.exists(): + srt_entries = parse_srt(str(srt_path)) + + repair_stats = await repair_tree( + index, empty_issues, vlm, llm, frames_dir, srt_entries + ) + stats["l3_repaired"] = repair_stats.l3_repaired + stats["l2_regenerated"] = repair_stats.l2_regenerated + stats["l1_regenerated"] = repair_stats.l1_regenerated + logger.info( + "[{}] 修复完成: L3={}, L2={}, L1={}", + vid, repair_stats.l3_repaired, repair_stats.l2_regenerated, + repair_stats.l1_regenerated, + ) + + # Step 3: 质量校验 + verify_stats = verify_tree(index) + total_removed = ( + verify_stats.l2_entities_removed + + verify_stats.l2_visible_text_removed + + verify_stats.l1_visible_text_removed + + verify_stats.l1_key_entities_removed + ) + stats["verify_removed"] = total_removed + if total_removed > 0: + logger.info("[{}] 校验删除 {} 项不可靠内容", vid, total_removed) + + # Step 4: Q&A 反向补全 + questions_path = questions_dir / f"{vid}.json" + if questions_path.exists(): + with open(questions_path, encoding="utf-8") as f: + questions = json.load(f) + if isinstance(questions, list) and questions: + logger.info("[{}] 开始 Q&A 补全 ({} 道题)...", vid, len(questions)) + srt_text = "" + srt_path = srt_dir / f"{vid}.srt" + if srt_path.exists(): + srt_text = srt_path.read_text(encoding="utf-8", errors="ignore") + + try: + supplement_stats = await supplement_tree( + index, questions, llm, srt_text=srt_text + ) + stats["facts_injected"] = supplement_stats.facts_injected + if supplement_stats.facts_injected > 0: + logger.info( + "[{}] 补全注入 {} 个事实", vid, supplement_stats.facts_injected + ) + except Exception as exc: + logger.error("[{}] Q&A 补全失败: {}", vid, exc) + logger.info("[{}] Q&A 补全完成", vid) + + # Step 5: 保存 + index.save_json(str(tree_path)) + logger.info("[{}] 已保存", vid) + + except Exception as exc: + stats["error"] = str(exc) + logger.error("[{}] 修复失败: {}", vid, exc) + + return stats + + +# --------------------------------------------------------------------------- +# 主流程 +# --------------------------------------------------------------------------- + + +async def main_async(args: argparse.Namespace) -> None: + """异步主流程:遍历所有视频,逐个修复。""" + videos_dir = Path(args.videos_dir) + srt_dir = Path(args.srt_dir) + questions_dir = Path(args.questions_dir) + + # 扫描所有视频 + vid_dirs = sorted( + d for d in videos_dir.iterdir() + if d.is_dir() and (d / "tree.json").exists() + ) + logger.info("发现 {} 个视频待修复", len(vid_dirs)) + + if args.dry_run: + logger.info("=== DRY RUN 模式:仅检测不修复 ===") + + # 构建客户端(dry_run 模式不需要) + llm, vlm = (None, None) if args.dry_run else _build_clients() + + # 逐视频修复 + all_stats = [] + start_time = time.time() + for idx, vid_dir in enumerate(vid_dirs): + vid = vid_dir.name + tree_path = vid_dir / "tree.json" + frames_dir = vid_dir # frame_path 已含 "frames/" 前缀,不再嵌套 + + logger.info( + "[{}/{}] 开始修复 {}", + idx + 1, len(vid_dirs), vid, + ) + + stats = await _repair_one_video( + vid, tree_path, frames_dir, srt_dir, questions_dir, + llm, vlm, dry_run=args.dry_run, + ) + all_stats.append(stats) + + # 每 10 个视频汇总一次 + if (idx + 1) % 10 == 0: + elapsed = time.time() - start_time + rate = (idx + 1) / elapsed * 60 + logger.info( + "进度: {}/{}, 已用 {:.0f}s, 速率 {:.1f} 视频/分钟", + idx + 1, len(vid_dirs), elapsed, rate, + ) + + # 最终汇总 + elapsed = time.time() - start_time + total_issues = sum(s["issues_found"] for s in all_stats) + total_repaired = sum(s["l3_repaired"] for s in all_stats) + total_injected = sum(s["facts_injected"] for s in all_stats) + total_errors = sum(1 for s in all_stats if s["error"]) + + logger.info("=" * 60) + logger.info("修复完成") + logger.info(" 视频总数: {}", len(all_stats)) + logger.info(" 问题总数: {}", total_issues) + logger.info(" L3 修复数: {}", total_repaired) + logger.info(" 事实注入数: {}", total_injected) + logger.info(" 失败数: {}", total_errors) + logger.info(" 总耗时: {:.0f}s", elapsed) + logger.info("=" * 60) + + if total_errors > 0: + logger.warning("以下视频修复失败:") + for s in all_stats: + if s["error"]: + logger.warning(" {}: {}", s["vid"], s["error"]) + + +def parse_args() -> argparse.Namespace: + """解析命令行参数。""" + parser = argparse.ArgumentParser(description="树修复管线") + parser.add_argument( + "--videos-dir", + default="store/videos", + help="视频目录(默认: store/videos)", + ) + parser.add_argument( + "--srt-dir", + default="data/Video-MME/subtitle", + help="SRT 字幕目录(默认: data/Video-MME/subtitle)", + ) + parser.add_argument( + "--questions-dir", + default="store/questions/benchmarks/Video-MME", + help="问题 JSON 目录(默认: store/questions/benchmarks/Video-MME)", + ) + parser.add_argument( + "--dry-run", + action="store_true", + help="仅检测不修复,不调用 VLM/LLM", + ) + return parser.parse_args() + + +def main() -> None: + """同步入口。""" + args = parse_args() + (PROJECT_ROOT / "logs").mkdir(exist_ok=True) + asyncio.run(main_async(args)) + + +if __name__ == "__main__": + main()