feat(tools): batch tree build orchestration with shared API semaphore
This commit is contained in:
@@ -5,9 +5,14 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
@@ -17,6 +22,8 @@ from tools.build_trees import ( # noqa: E402
|
||||
_discover_pending,
|
||||
_find_srt_entries,
|
||||
_tree_is_complete,
|
||||
load_progress,
|
||||
main_async,
|
||||
)
|
||||
|
||||
|
||||
@@ -106,3 +113,121 @@ class TestFindSrtEntries:
|
||||
def test_missing_returns_none(self, tmp_path: Path) -> None:
|
||||
"""无同名 .srt 返回 None。"""
|
||||
assert _find_srt_entries(tmp_path / "vid.mp4", tmp_path) is None
|
||||
|
||||
|
||||
# ── 编排集成(桩 builder,无真实视频/LLM)──────────────────────
|
||||
|
||||
|
||||
class _StubBuilder:
|
||||
"""记录并发与注入信号量的桩 builder。"""
|
||||
|
||||
instances: list[_StubBuilder] = []
|
||||
inflight = 0
|
||||
max_inflight = 0
|
||||
|
||||
def __init__(self, vlm, llm, config, *, api_semaphore=None) -> None:
|
||||
"""记录注入的 api_semaphore 并登记实例。"""
|
||||
self.api_semaphore = api_semaphore
|
||||
_StubBuilder.instances.append(self)
|
||||
|
||||
async def build_async(self, video_path: str, srt_entries=None) -> TreeIndex:
|
||||
"""模拟建树:短暂 sleep 并统计并发峰值,返回最小合法树。"""
|
||||
_StubBuilder.inflight += 1
|
||||
_StubBuilder.max_inflight = max(_StubBuilder.max_inflight, _StubBuilder.inflight)
|
||||
await asyncio.sleep(0.02)
|
||||
_StubBuilder.inflight -= 1
|
||||
|
||||
l1 = L1Node(
|
||||
id="x_L1_000",
|
||||
card=L1Card("s", "室内", [], [], [], [], "线性"),
|
||||
time_range=(0.0, 1.0),
|
||||
children=[],
|
||||
)
|
||||
return TreeIndex(metadata=IndexMeta(video_path, "video"), roots=[l1])
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def batch_env(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> dict:
|
||||
"""5 个假视频 + 桩 builder + 隔离的 progress 路径。"""
|
||||
import tools.build_trees as bt
|
||||
|
||||
_StubBuilder.instances = []
|
||||
_StubBuilder.inflight = 0
|
||||
_StubBuilder.max_inflight = 0
|
||||
monkeypatch.setattr(bt, "VideoTreeBuilder", _StubBuilder)
|
||||
monkeypatch.setattr(bt, "_build_clients", lambda api_concurrency: (None, None))
|
||||
|
||||
videos = tmp_path / "videos"
|
||||
videos.mkdir()
|
||||
for i in range(5):
|
||||
(videos / f"v{i}.mp4").write_bytes(b"")
|
||||
|
||||
return {
|
||||
"videos": videos,
|
||||
"out": tmp_path / "out",
|
||||
"progress": tmp_path / "build_progress.json",
|
||||
}
|
||||
|
||||
|
||||
def _make_args(env: dict, video_concurrency: int = 2, limit: int = 0) -> argparse.Namespace:
|
||||
"""构造 main_async 所需的 CLI 参数命名空间。"""
|
||||
return argparse.Namespace(
|
||||
videos_dir=str(env["videos"]),
|
||||
out_dir=str(env["out"]),
|
||||
srt_dir=str(env["videos"]),
|
||||
video_concurrency=video_concurrency,
|
||||
limit=limit,
|
||||
progress_path=str(env["progress"]),
|
||||
)
|
||||
|
||||
|
||||
class TestOrchestration:
|
||||
"""main_async 编排行为(桩 builder)。"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_video_concurrency_capped(self, batch_env: dict) -> None:
|
||||
"""同时在建视频数不得超过 video_concurrency。"""
|
||||
await main_async(_make_args(batch_env, video_concurrency=2))
|
||||
assert _StubBuilder.max_inflight <= 2
|
||||
assert len(_StubBuilder.instances) == 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shared_api_semaphore(self, batch_env: dict) -> None:
|
||||
"""全部 builder 实例共享同一个 API Semaphore 对象。"""
|
||||
await main_async(_make_args(batch_env))
|
||||
sems = {id(b.api_semaphore) for b in _StubBuilder.instances}
|
||||
assert len(sems) == 1
|
||||
assert _StubBuilder.instances[0].api_semaphore is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trees_saved_and_progress_recorded(self, batch_env: dict) -> None:
|
||||
"""每个视频产出 tree.json 且 progress 记录全部完成。"""
|
||||
await main_async(_make_args(batch_env))
|
||||
for i in range(5):
|
||||
assert (batch_env["out"] / f"v{i}" / "tree.json").exists()
|
||||
|
||||
assert load_progress(batch_env["progress"]) == {f"v{i}" for i in range(5)}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resume_skips_finished(self, batch_env: dict) -> None:
|
||||
"""第二次运行跳过全部已完成视频。"""
|
||||
await main_async(_make_args(batch_env))
|
||||
n_first = len(_StubBuilder.instances)
|
||||
await main_async(_make_args(batch_env))
|
||||
assert len(_StubBuilder.instances) == n_first # 无新建
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_partial_completion_resume(self, batch_env: dict) -> None:
|
||||
"""部分完成后重跑只建剩余视频(模拟中断后恢复)。"""
|
||||
batch_env["progress"].write_text(
|
||||
json.dumps({"finished_video_ids": ["v0", "v1"]}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
await main_async(_make_args(batch_env))
|
||||
assert len(_StubBuilder.instances) == 3 # 仅 v2/v3/v4
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_limit(self, batch_env: dict) -> None:
|
||||
"""--limit 2 只建前两个(烟测入口)。"""
|
||||
await main_async(_make_args(batch_env, limit=2))
|
||||
assert len(_StubBuilder.instances) == 2
|
||||
|
||||
Reference in New Issue
Block a user