Files
Video-Tree-TRM5/tests/unit/test_video_split_cli.py
T

385 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""video_split_cli 两阶段编排 CLI 单元测试。
覆盖:
- 配置解析正确 + 缺关键项 fail loudSystemExit);
- 指纹计算在 main 中被调用(diag_fingerprint 收到 prompt/model/git-sha 三分量);
- run_pipeline 两阶段按序触发(Phase 1 诊断 → Phase 2 build_split);
- McNemar 功效护栏 fail loud
- --dry-run 用假 deps 跑通装配、不真调 LLM(Phase 1 空错题早返回)。
不真跑全量诊断:诊断与 build_split 均以 monkeypatch / 假 deps 替换。
"""
from __future__ import annotations
import asyncio
from types import SimpleNamespace
import pytest
from app.harness import video_split_cli as cli
_RAW_OK = {
"video_split": {
"baseline_run_id": "infer_adhoc",
"n_trainval": 100,
"epsilon": 0.1,
"report_floor": 27,
"val_wrong_min": 20,
"val_ratio": 0.3,
"seed": 7,
"floor_k": {"Counting Problem": 3},
},
"diag": {"prompt_version": "diagnose_v1", "model": "deepseek-v4-pro"},
}
def _config(**overrides) -> cli.VideoSplitConfig:
"""构造一个可用 VideoSplitConfigoverrides 覆盖单字段。"""
base = {
"baseline_run_id": "infer_adhoc",
"n_trainval": 100,
"epsilon": 0.1,
"report_floor": 27,
"val_wrong_min": 0,
"val_ratio": 0.3,
"seed": 7,
"floor_k": {"Counting Problem": 3},
"prompt_version": "diagnose_v1",
"model": "deepseek-v4-pro",
}
base.update(overrides)
return cli.VideoSplitConfig(**base)
def test_parse_config_ok():
"""完整 yaml 解析为 VideoSplitConfig,各旋钮逐一还原。"""
cfg = cli.parse_config(_RAW_OK)
assert cfg.baseline_run_id == "infer_adhoc"
assert cfg.n_trainval == 100
assert cfg.epsilon == 0.1
assert cfg.report_floor == 27
assert cfg.val_wrong_min == 20
assert cfg.val_ratio == 0.3
assert cfg.seed == 7
assert cfg.floor_k == {"Counting Problem": 3}
assert cfg.prompt_version == "diagnose_v1"
assert cfg.model == "deepseek-v4-pro"
def test_parse_config_missing_video_split_key_fails_loud():
"""video_split 段缺关键项 → SystemExitP5 fail loud)。"""
raw = {"video_split": dict(_RAW_OK["video_split"]), "diag": dict(_RAW_OK["diag"])}
del raw["video_split"]["n_trainval"]
with pytest.raises(SystemExit):
cli.parse_config(raw)
def test_parse_config_missing_diag_section_fails_loud():
"""缺 diag 段 → SystemExit。"""
with pytest.raises(SystemExit):
cli.parse_config({"video_split": dict(_RAW_OK["video_split"])})
def test_load_config_missing_file_fails_loud(tmp_path):
"""config 文件不存在 → SystemExit。"""
with pytest.raises(SystemExit):
cli.load_config(tmp_path / "nope.yaml")
def test_check_mcnemar_power_below_threshold_fails_loud():
"""val 错题数 < 阈 → SystemExit(功效不足)。"""
q = SimpleNamespace(question_id="q1")
pools = SimpleNamespace(validation=[q], correctness={"q1": False})
with pytest.raises(SystemExit):
cli.check_mcnemar_power(pools, val_wrong_min=5)
def test_check_mcnemar_power_zero_threshold_skips():
"""val_wrong_min=0 → 不检查,返回实际错题数。"""
q = SimpleNamespace(question_id="q1")
pools = SimpleNamespace(validation=[q], correctness={"q1": True})
assert cli.check_mcnemar_power(pools, val_wrong_min=0) == 0
def test_run_pipeline_orders_three_phases(monkeypatch, tmp_path):
"""run_pipeline 先补 Phase 0 INFRA T0、再 Phase 1 诊断、后 Phase 2 build_split(按序)。"""
calls: list[str] = []
class _SpyStore:
def __init__(self):
self.t0_rows: list = []
def upsert(self, row):
calls.append("t0_upsert")
self.t0_rows.append(row)
async def fake_diag(**kwargs):
calls.append("diagnosis")
assert kwargs["diag_fingerprint"] == "fp"
assert kwargs["wrong_ids"] == ["q1"]
def fake_build_split(**kwargs):
calls.append("build_split")
assert kwargs["diag_fingerprint"] == "fp"
pools = SimpleNamespace(validation=[], correctness={})
return SimpleNamespace(pools=pools, manifest={}, assignment={})
monkeypatch.setattr(cli, "run_baseline_diagnosis", fake_diag)
monkeypatch.setattr(cli, "build_split", fake_build_split)
# 一条 INFRA 空预测错题 → Phase 0 应补一行 T0(在诊断/切分之前)。
canonical_preds = [
{
"question_id": "q_infra",
"video_id": "v9",
"task_type": "Counting Problem",
"prediction": "",
"answer": "A",
"stop_reason": "error",
"correct": False,
}
]
store = _SpyStore()
result = asyncio.run(
cli.run_pipeline(
config=_config(),
fingerprint="fp",
diagnosis_deps=object(),
signal_store=store,
wrong_ids=["q1"],
questions={},
canonical_preds=canonical_preds,
harness_db=tmp_path / "h.db",
questions_dir=tmp_path,
out_dir=tmp_path / "out",
generated_at="2026-07-15T00:00:00Z",
)
)
assert calls == ["t0_upsert", "diagnosis", "build_split"] # Phase 0 先于诊断与切分
assert len(store.t0_rows) == 1
assert store.t0_rows[0].tier == "T0" and store.t0_rows[0].infra is True
assert result.pools.validation == []
def test_select_diagnosable_wrong_ids_excludes_infra_and_correct():
"""可诊断错题筛选:排除 INFRA / 空预测 / 正确题,保留非空非 INFRA 错题(保序)。"""
preds = [
{"question_id": "ok", "stop_reason": "finished", "prediction": "B", "correct": True},
{"question_id": "wrong", "stop_reason": "finished", "prediction": "C", "correct": False},
{"question_id": "infra", "stop_reason": "error", "prediction": "", "correct": False},
{"question_id": "parse", "stop_reason": "parse_error", "prediction": "x", "correct": False},
{"question_id": "empty", "stop_reason": "finished", "prediction": "", "correct": False},
]
assert cli.select_diagnosable_wrong_ids(preds) == ["wrong"]
def test_persist_infra_t0_rows_persists_only_infra_or_empty_wrong(tmp_path):
"""INFRA / 空预测错题落 T0infra=Trueerror_type/target=None);正确 / 可诊断错题不落。"""
from adapters.baseline_diagnosis_store import SqliteDiagnosisSignalStore
preds = [
{
"question_id": "ok",
"video_id": "v1",
"task_type": "Counting Problem",
"prediction": "A",
"answer": "A",
"stop_reason": "finished",
"correct": True,
},
{
"question_id": "diag_wrong",
"video_id": "v2",
"task_type": "Counting Problem",
"prediction": "C",
"answer": "A",
"stop_reason": "finished",
"correct": False,
},
{
"question_id": "infra_err",
"video_id": "v3",
"task_type": "OCR Problems",
"prediction": "",
"answer": "A",
"stop_reason": "error",
"correct": False,
},
{
"question_id": "parse_err",
"video_id": "v4",
"task_type": "Counting Problem",
"prediction": "",
"answer": "A",
"stop_reason": "parse_error",
"correct": False,
},
]
store = SqliteDiagnosisSignalStore(str(tmp_path / "h.db"))
n = cli.persist_infra_t0_rows(store, preds, "infer_adhoc", "fp")
assert n == 2 # 只有两条 INFRA 空预测错题
rows = {r.question_id: r for r in store.load("infer_adhoc", "fp")}
assert set(rows) == {"infra_err", "parse_err"}
for r in rows.values():
assert r.tier == "T0" and r.infra is True
assert r.error_type is None and r.evolution_target is None and r.cause_category is None
assert r.degraded is False
# 幂等:重复调用同 PK 覆盖,行数不变。
assert cli.persist_infra_t0_rows(store, preds, "infer_adhoc", "fp") == 2
assert len(store.load("infer_adhoc", "fp")) == 2
store.close()
def test_dry_run_computes_fingerprint_without_llm(monkeypatch, tmp_path, capsys):
"""--dry-rundiag_fingerprint 被调用、Phase 1 走空错题早返回、不真调 LLM。"""
fp_calls: list[tuple[str, str, str]] = []
def spy_fingerprint(prompt_version, model, code_version):
fp_calls.append((prompt_version, model, code_version))
return "dryfp"
monkeypatch.setattr(cli, "diag_fingerprint", spy_fingerprint)
monkeypatch.setattr(cli, "git_short_sha", lambda: "abc123")
# 用真实 config/video_split.yamlout-dir 指向 tmpdry-run 会建临时信号库再清理)。
cli.main(["--dry-run", "--config", "config/video_split.yaml", "--out-dir", str(tmp_path)])
assert len(fp_calls) == 1
prompt_version, model, sha = fp_calls[0]
assert sha == "abc123"
assert prompt_version == "diagnose_v1"
assert model == "deepseek-v4-pro"
# dry-run 不留下临时信号库
assert not (tmp_path / "_dry_run_signals.db").exists()
def test_build_diagnosis_deps_model_mismatch_fails_loud(monkeypatch, tmp_path):
"""config.diag.model 与 .env SEARCH_LLM_MODEL 不一致 → fail loud(指纹漂移防护)。"""
class _FakeSettings:
search_llm_model = "actual-model-in-env"
search_llm_base_url = "https://api.example"
search_llm_api_key = "sk-xxx"
monkeypatch.setattr(cli, "_DiagLLMSettings", lambda: _FakeSettings())
with pytest.raises(SystemExit) as exc:
cli.build_diagnosis_deps(
harness_db=tmp_path / "h.db",
store_dir=tmp_path,
video_ids=[],
concurrency=2,
expected_model="deepseek-v4-pro", # 与 env 不一致
)
# 报错须同时暴露两个值,便于人对齐
msg = str(exc.value)
assert "deepseek-v4-pro" in msg
assert "actual-model-in-env" in msg
def test_build_diagnosis_deps_missing_credentials_fails_loud(monkeypatch, tmp_path):
""".env 缺 search LLM 凭证 → fail loud(先于模型一致性校验)。"""
class _EmptySettings:
search_llm_model = ""
search_llm_base_url = ""
search_llm_api_key = ""
monkeypatch.setattr(cli, "_DiagLLMSettings", lambda: _EmptySettings())
with pytest.raises(SystemExit):
cli.build_diagnosis_deps(
harness_db=tmp_path / "h.db",
store_dir=tmp_path,
video_ids=[],
concurrency=2,
expected_model="deepseek-v4-pro",
)
def test_git_short_sha_nonempty():
"""仓库内 git_short_sha 返回非空短 SHA。"""
sha = cli.git_short_sha()
assert sha
assert len(sha) >= 4
def test_build_split_refuses_overwrite_without_force(tmp_path):
"""已存在指纹不同的 pools.json 时,force=False 必须报错不覆盖。"""
from app.harness.build_split import _guard_frozen_products
out_path = tmp_path / "pools.json"
out_path.write_text('{"split_mode":"global"}', encoding="utf-8")
manifest_path = tmp_path / "split_manifest.json"
with pytest.raises(FileExistsError, match="已存在冻结产物"):
_guard_frozen_products(out_path, manifest_path, force=False)
def test_build_split_force_backs_up_old(tmp_path):
"""force=True 时旧产物被备份为 .bak.* 再允许覆盖。"""
from app.harness.build_split import _guard_frozen_products
out_path = tmp_path / "pools.json"
out_path.write_text('{"old":1}', encoding="utf-8")
manifest_path = tmp_path / "split_manifest.json"
manifest_path.write_text('{"pools_sha256":"deadbeef00000000"}', encoding="utf-8")
_guard_frozen_products(out_path, manifest_path, force=True)
baks = list(tmp_path.glob("pools.json.bak.*"))
assert len(baks) == 1, f"未备份旧产物: {list(tmp_path.iterdir())}"
def test_build_split_force_twice_same_suffix_keeps_both(tmp_path):
"""连续两次 force 命中同后缀时,第二次备份不覆盖第一次(追加递增序号)。"""
from app.harness.build_split import _guard_frozen_products
out_path = tmp_path / "pools.json"
manifest_path = tmp_path / "split_manifest.json"
# 第一次 force:两文件都无 manifest 里的 pools_sha256 之外的差异,后缀恒为同值。
out_path.write_text('{"gen":1}', encoding="utf-8")
manifest_path.write_text('{"pools_sha256":"deadbeef00000000"}', encoding="utf-8")
_guard_frozen_products(out_path, manifest_path, force=True)
# 第二次 force:写入相同 sha 的新产物,触发同后缀备份。
out_path.write_text('{"gen":2}', encoding="utf-8")
manifest_path.write_text('{"pools_sha256":"deadbeef00000000"}', encoding="utf-8")
_guard_frozen_products(out_path, manifest_path, force=True)
pools_baks = sorted(p.name for p in tmp_path.glob("pools.json.bak.*"))
assert len(pools_baks) == 2, f"同后缀第二次备份覆盖了第一次: {pools_baks}"
def test_build_split_force_rollback_on_partial_failure(tmp_path, monkeypatch):
"""备份第二个文件失败时,第一个已备份文件被 rollback 回原名(原子性)。"""
from pathlib import Path
from app.harness import build_split as bs
out_path = tmp_path / "pools.json"
manifest_path = tmp_path / "split_manifest.json"
out_path.write_text('{"gen":1}', encoding="utf-8")
manifest_path.write_text('{"pools_sha256":"deadbeef00000000"}', encoding="utf-8")
real_rename = Path.rename
calls = {"n": 0}
def flaky_rename(self, target):
# 第一次 rename(备份 pools.json)成功,第二次(备份 manifest)抛错。
calls["n"] += 1
if calls["n"] == 2:
raise OSError("模拟第二个备份 rename 失败")
return real_rename(self, target)
monkeypatch.setattr(Path, "rename", flaky_rename)
with pytest.raises(OSError, match="模拟第二个备份"):
bs._guard_frozen_products(out_path, manifest_path, force=True)
# rollback 后:两原文件仍在原位,无残留 .bak.*
assert out_path.exists(), "第一个文件未被 rollback 回原名"
assert manifest_path.exists(), "第二个文件不应被移动"
assert not list(tmp_path.glob("*.bak.*")), f"残留半备份: {list(tmp_path.iterdir())}"