muse-agent-example/tests/skills/写下一章/test_dispatch_writer_bridge.py

181 lines
7.0 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.

#!/usr/bin/env python3
"""生产写手框架派发桥的离线测试(不连真实框架、不连真实模型)。
固定合同:任务包装配(角色/冻结输入/探索白名单)、会话按章稳定、
成功路径绑定候选信封与证据引用、派发失败与证据缺失一律失败关闭。
"""
from __future__ import annotations
import json
import pathlib
import sys
import tempfile
import unittest
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
SCRIPT_DIR = PROJECT_ROOT / "muse" / "content" / "work" / "skills" / "generate" / "写下一章" / "scripts"
DISPATCH_TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "派发智能体任务"
CHECK_TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "核对内容一致性"
WNC_TEST_DIR = PROJECT_ROOT / "tests" / "skills" / "写下一章"
for path in (SCRIPT_DIR, DISPATCH_TEST_DIR, CHECK_TEST_DIR, WNC_TEST_DIR):
if str(path) not in sys.path:
sys.path.insert(0, str(path))
from dispatch_agent_task import DEFAULT_RUN_DIR_ROOT # noqa: E402
from dispatch_writer_bridge import ( # noqa: E402
DispatchWriterError,
build_writer_dispatch_spec,
run_writer_via_dispatch,
writer_session_paths,
)
from read_tools import TOOL_REGISTRY # noqa: E402
from test_check_writer_candidate import _valid_pair # noqa: E402
from test_dispatch_agent_task import ( # noqa: E402
FIXED_MODEL,
RecordingConnect,
fake_launcher,
pi_stream_lines,
)
class BridgeConnect(RecordingConnect):
"""派发桥专用假连接:补证证据引用查询返回 raw_row(可置 None 模拟缺证)。"""
def __init__(self):
super().__init__()
self.raw_row = (9001, 7777)
def __call__(self, *args, **kwargs):
outer = self
class _Cursor:
def execute(self, sql, params=None):
outer.log.append((sql, params))
return self
def fetchone(self):
sql = outer.log[-1][0]
if sql.startswith("SELECT run_id, work_id"):
return ("row", None, None, "running")
if sql.startswith("SELECT id, raw_content_id"):
return outer.raw_row
if sql.startswith("INSERT INTO example_run"):
return ("row",)
if sql.startswith("UPDATE example_run"):
state = "failed" if "failed" in (outer.log[-1][1] or ()) else "completed"
return ("row", state, None)
RecordingConnect.next_id += 1
return (RecordingConnect.next_id,)
def fetchall(self):
return []
def commit(self):
outer.log.append(("COMMIT", None))
def rollback(self):
outer.log.append(("ROLLBACK", None))
class _Ctx:
def __enter__(self):
return _Cursor()
def __exit__(self, *exc):
return False
return _Ctx()
def _candidate_body() -> str:
return "茧撕开舱门的一瞬,林深听见了深渊的回声。" * 8
class SpecBuildTest(unittest.TestCase):
def test_spec_freezes_creative_input_and_read_tools(self):
context, _ = _valid_pair()
spec = build_writer_dispatch_spec(
context, target_chapter=3, human_instruction="让恐惧落在身体上", candidate_version=2,
)
self.assertEqual(spec["role"], "writer")
self.assertEqual(spec["input"]["workId"], context["workId"])
self.assertEqual(spec["input"]["targetChapter"], 3)
self.assertEqual(spec["input"]["candidateVersion"], 2)
self.assertIn("让恐惧落在身体上", spec["taskPrompt"])
self.assertIn("creativeInput", spec["input"])
# 探索白名单 = 工具 server 登记表(单一事实源,防漂移)。
self.assertEqual(spec["toolAllowlist"], sorted(TOOL_REGISTRY))
self.assertEqual(
spec["outputSchema"]["required"], ["candidateBody"])
def test_session_paths_are_stable_per_chapter(self):
sid_a, dir_a = writer_session_paths(12, 3)
sid_b, dir_b = writer_session_paths(12, 3)
sid_c, dir_c = writer_session_paths(12, 4)
self.assertEqual((sid_a, dir_a), (sid_b, dir_b))
self.assertNotEqual(sid_a, sid_c)
self.assertNotEqual(dir_a, dir_c)
class DispatchRoundTripTest(unittest.TestCase):
def setUp(self):
self.tmp = pathlib.Path(tempfile.mkdtemp())
self.context, _ = _valid_pair()
# 运行目录按派发运行号固定;测试前清理,避免审计根冲突。
self.dispatch_run_id = f"{self.context['runId']}-writer-v1"
import shutil
shutil.rmtree(DEFAULT_RUN_DIR_ROOT / self.dispatch_run_id, ignore_errors=True)
def _run(self, lines, connect):
return run_writer_via_dispatch(
self.context,
candidate_version=1,
repo_root=PROJECT_ROOT,
provider="p",
model=FIXED_MODEL,
thinking="low",
human_instruction="",
spec_path=self.tmp / "task.json",
launcher=fake_launcher(lines),
connect_factory=connect,
sqlite_path=self.tmp / "ledger.db",
)
def test_success_binds_envelope_and_evidence_ref(self):
final_text = json.dumps({"candidateBody": _candidate_body()}, ensure_ascii=False)
connect = BridgeConnect()
envelope, receipt, raw_ref = self._run(
pi_stream_lines(final_text, model=FIXED_MODEL), connect
)
# 身份、哈希、版本由桥绑定,不来自模型。
self.assertEqual(envelope["runId"], self.context["runId"])
self.assertEqual(envelope["candidateVersion"], 1)
self.assertTrue(envelope["candidateSha256"].startswith("sha256:"))
self.assertEqual(envelope["candidateBody"], _candidate_body())
# 回执适配供生产账本消费;证据引用来自派发运行的调用账。
self.assertEqual(receipt.requested_model_id, f"p/{FIXED_MODEL}")
self.assertTrue(receipt.dispatch_run_id.startswith(self.context["runId"]))
self.assertEqual(raw_ref, (9001, 7777))
# 任务包落盘可审计。
saved_spec = json.loads((self.tmp / "task.json").read_text(encoding="utf-8"))
self.assertEqual(saved_spec["role"], "writer")
def test_dispatch_failure_fails_closed(self):
connect = BridgeConnect()
with self.assertRaises(DispatchWriterError) as caught:
# 空事件流 -> 框架层失败关闭。
self._run([], connect)
self.assertNotEqual(caught.exception.code, "")
def test_missing_raw_reference_fails_closed(self):
final_text = json.dumps({"candidateBody": _candidate_body()}, ensure_ascii=False)
connect = BridgeConnect()
connect.raw_row = None
with self.assertRaises(DispatchWriterError) as caught:
self._run(pi_stream_lines(final_text, model=FIXED_MODEL), connect)
self.assertEqual(caught.exception.code, "DISPATCH_EVIDENCE_MISSING")
if __name__ == "__main__":
unittest.main()