150 lines
5.0 KiB
Python
150 lines
5.0 KiB
Python
#!/usr/bin/env python3
|
||
"""writer 统一落库适配器的输入合同离线测试。"""
|
||
import pathlib
|
||
import sys
|
||
import unittest
|
||
|
||
|
||
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
|
||
SCRIPT_DIR = PROJECT_ROOT / "muse" / "content" / "work" / "skills" / "generate" / "write-next-chapter" / "scripts"
|
||
sys.path.insert(0, str(SCRIPT_DIR))
|
||
|
||
import persist_writer_run as writer_persist # noqa: E402
|
||
|
||
|
||
class WriterPersistenceContractTest(unittest.TestCase):
|
||
def test带前缀哈希归一(self):
|
||
value = "sha256:" + "a" * 64
|
||
self.assertEqual(writer_persist._bare_hash(value, "x"), "a" * 64)
|
||
|
||
def test非法哈希失败关闭(self):
|
||
with self.assertRaises(writer_persist.WriterPersistenceError):
|
||
writer_persist._bare_hash("not-a-hash", "candidateSha256")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|
||
|
||
|
||
class _FakeCursor:
|
||
"""按 SQL 关键词回放固定行的假游标(只验证落库合同,不连库)。"""
|
||
|
||
def __init__(self, log):
|
||
self._log = log
|
||
self._last = None
|
||
self._ids = iter(range(100, 1000))
|
||
|
||
def execute(self, sql, params=None):
|
||
self._log.append((sql, params))
|
||
self._last = sql
|
||
return self
|
||
|
||
def fetchone(self):
|
||
sql = self._last or ""
|
||
if sql.startswith("SELECT id,candidate_sha256,state"):
|
||
return None
|
||
if "COALESCE(MAX(revision)" in sql:
|
||
return (0,)
|
||
if "RETURNING id" in sql:
|
||
return (next(self._ids),)
|
||
return (next(self._ids),)
|
||
|
||
def commit(self):
|
||
self._log.append(("COMMIT", None))
|
||
|
||
def rollback(self):
|
||
self._log.append(("ROLLBACK", None))
|
||
|
||
|
||
class _FakeConnCtx:
|
||
def __init__(self, log):
|
||
self._log = log
|
||
|
||
def __enter__(self):
|
||
return _FakeCursor(self._log)
|
||
|
||
def __exit__(self, *exc):
|
||
return False
|
||
|
||
|
||
def _minimal_inputs():
|
||
context = {
|
||
"runId": "run-persist-dispatch-1",
|
||
"workId": 12,
|
||
"targetChapter": 3,
|
||
"attempt": 1,
|
||
"contextSnapshot": {"contextSha256": "sha256:" + "b" * 64},
|
||
}
|
||
candidate = {
|
||
"runId": "run-persist-dispatch-1",
|
||
"attempt": 1,
|
||
"candidateVersion": 1,
|
||
"candidateSha256": "sha256:" + "c" * 64,
|
||
"candidateBody": "正文候选",
|
||
}
|
||
mechanical = {"passed": True, "requirements": []}
|
||
return context, candidate, mechanical
|
||
|
||
|
||
class DispatchRawRefTest(unittest.TestCase):
|
||
"""阶段 E:派发链模型证据在派发运行下,编排显式传引用,禁止两套记账。"""
|
||
|
||
def _patch_runtime(self):
|
||
import contextlib
|
||
|
||
log = []
|
||
|
||
@contextlib.contextmanager
|
||
def _patches():
|
||
orig_connect = writer_persist.connect
|
||
orig_start = writer_persist.start_run
|
||
orig_finish = writer_persist.finish_run
|
||
orig_lesson = writer_persist._propose_writer_lesson
|
||
writer_persist.connect = lambda *a, **k: _FakeConnCtx(log)
|
||
writer_persist.start_run = lambda *a, **k: None
|
||
writer_persist.finish_run = lambda *a, **k: None
|
||
writer_persist._propose_writer_lesson = lambda *a, **k: None
|
||
try:
|
||
yield log
|
||
finally:
|
||
writer_persist.connect = orig_connect
|
||
writer_persist.start_run = orig_start
|
||
writer_persist.finish_run = orig_finish
|
||
writer_persist._propose_writer_lesson = orig_lesson
|
||
|
||
return _patches()
|
||
|
||
def test显式引用绕过本运行查询并落回执(self):
|
||
context, candidate, mechanical = _minimal_inputs()
|
||
|
||
def _forbid_raw_lookup(*args, **kwargs):
|
||
raise AssertionError("显式引用模式下不得查询本运行调用账")
|
||
|
||
with self._patch_runtime() as log:
|
||
orig_raw = writer_persist._raw_response
|
||
writer_persist._raw_response = _forbid_raw_lookup
|
||
try:
|
||
result = writer_persist.persist_writer_execution(
|
||
context, candidate, receipt=None, mechanical_report=mechanical,
|
||
semantic_report=None, assemble_result=None,
|
||
writer_raw_ref=(5001, 6002),
|
||
)
|
||
finally:
|
||
writer_persist._raw_response = orig_raw
|
||
self.assertEqual(result["status"], "persisted")
|
||
self.assertEqual(result["raw_content_id"], 6002)
|
||
# 候选与回执都落了库。
|
||
inserts = [sql for sql, _ in log if sql.startswith("INSERT INTO")]
|
||
self.assertTrue(any("example_candidate" in sql for sql in inserts))
|
||
self.assertTrue(any("example_run_receipt" in sql for sql in inserts))
|
||
|
||
def test引用缺项失败关闭(self):
|
||
context, candidate, mechanical = _minimal_inputs()
|
||
with self._patch_runtime():
|
||
with self.assertRaises(writer_persist.WriterPersistenceError):
|
||
writer_persist.persist_writer_execution(
|
||
context, candidate, receipt=None, mechanical_report=mechanical,
|
||
semantic_report=None, assemble_result=None,
|
||
writer_raw_ref=(None, None),
|
||
)
|