diff --git a/nodes/llm_filter.py b/nodes/llm_filter.py index 9e544c6..2a57713 100644 --- a/nodes/llm_filter.py +++ b/nodes/llm_filter.py @@ -33,6 +33,7 @@ from __future__ import annotations import json import os import re +import threading import time import urllib.error import urllib.request @@ -71,6 +72,39 @@ _ALL_CATEGORIES = ( ) DELETE_CATEGORIES = {CATEGORY_GARBAGE, CATEGORY_OVERLAY, CATEGORY_NOISE} +# 判定存档文件名:位于节点 output_dir,每行 {"index": 条目标引, "category": 类别}。 +# 每条 LLM 判定成功即追加一行;进程被杀/节点失败(如 429 限流)后重跑时, +# 只对未判定的条目重新调用 LLM,已判定结果直接复用(类似 OCR 的断点存档)。 +_PARTIAL_NAME = "filter_partial.jsonl" + +# 判定存档追加写锁:多线程判定并发完成时串行化追加,避免行交错。 +_partial_lock = threading.Lock() + + +def _load_partial(output_dir: Path) -> dict[int, str]: + """读取判定存档,返回 {条目标引: 类别};无存档/损坏行跳过。""" + path = output_dir / _PARTIAL_NAME + if not path.is_file(): + return {} + result: dict[int, str] = {} + for line in path.read_text(encoding="utf-8").splitlines(): + if not line.strip(): + continue + try: + item = json.loads(line) + except json.JSONDecodeError: + # 进程被杀时可能残留半行写入:跳过,对应条目视为未判定。 + continue + result[int(item["index"])] = str(item["category"]) + return result + + +def _append_partial(output_dir: Path, index: int, category: str) -> None: + """线程安全地把一条判定结果追加到存档(成功判定后立即落盘)。""" + with _partial_lock: + with (output_dir / _PARTIAL_NAME).open("a", encoding="utf-8") as fh: + fh.write(json.dumps({"index": index, "category": category}, ensure_ascii=False) + "\n") + # 规则层正则:横线装饰(含全角/半角横线、下划线、中点、句点等符号组合)。 _DASH_RE = re.compile(r"^[\s\-—_~=•・。..、]+$") # 规则层正则:URL / 邮箱。 @@ -295,7 +329,11 @@ def invoke(request: InvokeRequest) -> InvokeResponse: llm_needed = [i for i, verdict in enumerate(rule_verdicts) if verdict is None] # 阶段 2:LLM 分类层(去重:相同文本只判一次,上下文取首次出现)。 - cat_by_index: dict[int, str] = {} + # 节点级断点存档:output_dir 提前建好,重跑时只重判未判定条目。 + output_dir = Path(request.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + partial = _load_partial(output_dir) + cat_by_index: dict[int, str] = dict(partial) if llm_needed: if dedupe: first_of_key: dict[str, int] = {} @@ -304,20 +342,24 @@ def invoke(request: InvokeRequest) -> InvokeResponse: key = _dedup_key(entries[i]["text"]) if key not in first_of_key: first_of_key[key] = i - pool_indices.append(i) + # 断点续跑:该键首次出现已在存档判定过则跳过(结果复用)。 + if i not in partial: + pool_indices.append(i) else: - pool_indices = llm_needed - + # 断点续跑:只处理未判定的条目。 + pool_indices = [i for i in llm_needed if i not in partial] # 单条判断的工作函数:返回类别词;overlay_tokens 用于上下文净化。 def judge_one(index: int) -> str: try: - return _judge_category(entries, index, context_size, request.params, overlay_tokens) + category = _judge_category(entries, index, context_size, request.params, overlay_tokens) except urllib.error.HTTPError as exc: # 限流/服务端错误:通知线程池临时降低最大并发,避免持续超配额。 if exc.code == 429 or 500 <= exc.code < 600: pool.report_failure() raise - + # 判定成功立即落盘(断点存档):失败/中断后重跑不重复调用已判定条目。 + _append_partial(output_dir, index, category) + return category # 进度日志:打印已判定条数、总数、平均处理速度(条/s)、最近窗口 # 平均单条耗时与当前线程数(与 OCR 节点同一回调协议)。 def log_progress(done: int, total: int, rate: float, avg_time: float, workers: int) -> None: diff --git a/scripts/rerun_filter_011d01f19999.py b/scripts/rerun_filter_011d01f19999.py new file mode 100644 index 0000000..13a2bf2 --- /dev/null +++ b/scripts/rerun_filter_011d01f19999.py @@ -0,0 +1,26 @@ +"""一次性脚本:用上下文净化后的 llm_filter 重跑 run_011d01f19999 的 filter 节点。 + +目的:验证"上下文净化"修复对误删真实对话的恢复效果(对比原输出 885 保留 / 805 删除)。 +""" +from pathlib import Path +from dotenv import load_dotenv + +load_dotenv(Path(".env")) + +from nodes.llm_filter import invoke # noqa: E402 +from wov_sdk.models import InvokeRequest # noqa: E402 + +run_root = Path("data/storage/runs/run_011d01f19999/steps") +out_dir = run_root / "filter_rerun_v2" + +resp = invoke( + InvokeRequest( + run_id="run_011d01f19999", + node_instance_id="", + inputs={"srt_uri": str(run_root / "ocr/subtitle.srt")}, + params={"pool_min_workers": 8, "pool_max_workers": 8}, + output_dir=str(out_dir), + ) +) +print("status:", resp.status, "| error:", resp.error) +print("outputs:", resp.outputs) diff --git a/tests/test_llm_filter.py b/tests/test_llm_filter.py index 3f319b5..64844f1 100644 --- a/tests/test_llm_filter.py +++ b/tests/test_llm_filter.py @@ -709,3 +709,111 @@ def test_real_run_rules_and_dialogue_regression(monkeypatch, tmp_path) -> None: }) assert len(fake.bodies) == unique_llm assert len(fake.bodies) < len(entries) # 去重确实省调用。 + + +class _FailOnTarget429: + """模拟持续限流:目标字幕含指定词时恒抛 429,其余正常返回 dialogue。 + + 用于构造"部分条目成功、个别条目持续 429"的断点重跑场景:第一次 invoke + 整体失败但成功条目已写存档;解除限流后第二次 invoke 只重判失败条目。 + """ + + def __init__(self, marker: str) -> None: + self.marker = marker + self.enabled = True + self.bodies: list[dict] = [] + + def __call__(self, request, timeout=None): + body = json.loads(request.data.decode("utf-8")) + self.bodies.append(body) + target = next( + line for line in body["messages"][1]["content"].splitlines() + if line.startswith("【目标】") + ) + if self.enabled and self.marker in target: + raise urllib.error.HTTPError(request.full_url, 429, "rate limited", {}, None) + payload = json.dumps({"choices": [{"message": {"content": "dialogue"}}]}).encode() + return FakeResponse(payload) + + +def _llm_invoke(srt_text: str, out: Path) -> tuple[object, Path]: + """用给定 SRT 文本构造并执行一次 llm-filter invoke,返回 (响应, 输入文件)。""" + srt = out.parent / "in.srt" + srt.write_text(srt_text, encoding="utf-8") + return ( + invoke( + InvokeRequest( + run_id="r", node_instance_id="", + inputs={"srt_uri": str(srt)}, + params={}, + output_dir=str(out), + ) + ), + srt, + ) + + +def test_load_and_append_partial(tmp_path) -> None: + """判定存档读写:无存档/损坏行跳过,追加后可读回。""" + from nodes.llm_filter import _PARTIAL_NAME, _append_partial, _load_partial + + out = tmp_path / "out" + assert _load_partial(out) == {} # 目录不存在 → 空。 + out.mkdir() + assert _load_partial(out) == {} # 无存档 → 空。 + # 损坏行(半行写入)跳过,正常行读回。 + (out / _PARTIAL_NAME).write_text( + '{"index": 0, "category": "dialogue"}\n\n{broken\n{"index": 3, "category": "garbage"}\n', + encoding="utf-8", + ) + assert _load_partial(out) == {0: "dialogue", 3: "garbage"} + # 追加一条后读回。 + _append_partial(out, 5, "overlay") + assert _load_partial(out)[5] == "overlay" + + +def test_invoke_resume_skips_archived_judgments(monkeypatch, tmp_path) -> None: + """断点存档:已判定条目重跑时不重复调用 LLM(去重后只补判未判定)。""" + fake = _patch_llm(monkeypatch, contents=["dialogue"] * 4) + out = tmp_path / "out" + out.mkdir() + # 预置存档:索引 0、1 已判定(模拟上次失败前已完成的部分)。 + from nodes.llm_filter import _append_partial + + _append_partial(out, 0, "dialogue") + _append_partial(out, 1, "dialogue") + response, _srt = _llm_invoke(_SRT, out) + assert response.status == "completed", response.error + # 4 条唯一文本中 2 条已存档,只调用剩余 2 条。 + assert len(fake.bodies) == 2 + # 输出与全量判定一致:全部 dialogue → 4 条都保留。 + assert response.outputs["kept"] == 4 + assert Path(response.outputs["srt_uri"]).read_text(encoding="utf-8").count("-->") == 4 + + +def test_invoke_partial_failure_then_resume(monkeypatch, tmp_path) -> None: + """真实断点重跑:个别条目持续 429 → 整体失败但成功判定已写存档, + 解除限流后重跑只补判失败条目,最终产物与一次跑完一致。""" + fake = _FailOnTarget429(marker="无意义") + monkeypatch.setattr("nodes.llm_filter.urllib.request.urlopen", fake) + monkeypatch.setattr("nodes.llm_filter.time.sleep", lambda s: None) + out = tmp_path / "out" + first, _srt = _llm_invoke(_SRT, out) + assert first.status == "failed" + assert "429" in (first.error or "") + calls_first = len(fake.bodies) + # 成功判定的 3 条已写存档;持续 429 的"答:无意义杂项"(索引 2)不在存档。 + from nodes.llm_filter import _load_partial + + partial = _load_partial(out) + assert 2 not in partial + assert len(partial) == 3 + # 解除限流后重跑:只补判失败条目(1 次调用),其余复用存档。 + fake.enabled = False + second, _srt = _llm_invoke(_SRT, out) + assert second.status == "completed", second.error + assert len(fake.bodies) == calls_first + 1 + # 输出:解除限流后全部判定为 dialogue → 4 条都保留(与"一次跑完"一致)。 + out_text = Path(second.outputs["srt_uri"]).read_text(encoding="utf-8") + assert out_text.count("-->") == 4 + assert "00:00:09,000" in out_text