353 lines
14 KiB
Python
353 lines
14 KiB
Python
#!/usr/bin/env python3
|
||
"""无工具正文写手 adapter 的派发合同、合同绑定与失败关闭测试。
|
||
|
||
执行器是 muse_role 的无 CLI 固定模型策略:身份提示与中心角色合同注入系统提示词、冻结输入作
|
||
prompt、输出按 schema 校验。测试注入与 muse_llm.chat_governed 同签名的假实现,
|
||
不触发真实模型。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import copy
|
||
import hashlib
|
||
import json
|
||
import pathlib
|
||
import sys
|
||
import unittest
|
||
from decimal import Decimal
|
||
|
||
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
|
||
SCRIPT_DIR = PROJECT_ROOT / ".agent" / "skills" / "write-next-chapter" / "scripts"
|
||
READ_CONTEXT_DIR = PROJECT_ROOT / ".agent" / "skills" / "assemble-context" / "scripts"
|
||
TEST_CONTEXT_DIR = PROJECT_ROOT / "tests" / "skills" / "assemble-context"
|
||
sys.path.insert(0, str(SCRIPT_DIR))
|
||
sys.path.insert(0, str(READ_CONTEXT_DIR))
|
||
sys.path.insert(0, str(TEST_CONTEXT_DIR))
|
||
|
||
from test_writer_contract import valid_context, valid_draft # noqa: E402
|
||
from writer_contract import build_writer_creative_input, canonical_json, retrieval_identity # noqa: E402
|
||
from muse_role import ( # noqa: E402
|
||
FIXED_OPUS_MODEL_ID,
|
||
FIXED_OPUS_POLICY_ALIAS,
|
||
ROLE_TASK_SEPARATOR,
|
||
build_dispatch_system_prompt,
|
||
)
|
||
from run_writer import ( # noqa: E402
|
||
WRITER_ROLE_PROMPT,
|
||
WriterAdapterError,
|
||
build_production_length_contracts,
|
||
build_writer_execution_profile,
|
||
calculate_dynamic_output_contract,
|
||
run_writer,
|
||
)
|
||
|
||
USAGE = {"input_tokens": 100, "output_tokens": 20}
|
||
ACTUAL_MODEL = FIXED_OPUS_MODEL_ID
|
||
|
||
|
||
def _bound_context() -> dict:
|
||
"""构造带一条事实证据且可供 adapter 绑定检查的合法上下文。"""
|
||
|
||
context = valid_context()
|
||
context["factEvidence"] = [
|
||
{
|
||
"evidenceId": "fact-1",
|
||
"fact": "林澈仍在圣蒂曼城内",
|
||
"sourceType": "canonical_state",
|
||
"sourceRef": {
|
||
"sourceId": "state:8",
|
||
"sourceVersion": "state-v1",
|
||
"sourceType": "canonical_state",
|
||
},
|
||
"contentSha256": "sha256:" + "a" * 64,
|
||
"riskLevel": "high",
|
||
}
|
||
]
|
||
context["contextSnapshot"]["contextSha256"] = retrieval_identity(context)
|
||
return context
|
||
|
||
|
||
def _writer_draft(_context: dict, *, body: str | None = None) -> dict:
|
||
"""构造 writer 模型允许返回的单字段草稿。"""
|
||
|
||
return valid_draft(body=body if body is not None else "文" * 4000)
|
||
|
||
|
||
class FakeChat:
|
||
"""与 muse_llm.chat_governed 同签名的假实现,确保测试不调用真实模型。
|
||
|
||
payload 为字符串时按成功返回;None 模拟治理链全链耗尽;异常实例模拟底座失败。
|
||
"""
|
||
|
||
def __init__(self, payload):
|
||
self.payload = payload
|
||
self.calls: list[dict] = []
|
||
|
||
def __call__(self, prompt, system=None, temperature=0.2, top_p=None, *,
|
||
run_id=None, caller=None, persist_call=None, **_kwargs):
|
||
self.calls.append(
|
||
{
|
||
"prompt": prompt,
|
||
"system": system,
|
||
"temperature": temperature,
|
||
"run_id": run_id,
|
||
"caller": caller,
|
||
}
|
||
)
|
||
if isinstance(self.payload, BaseException):
|
||
raise self.payload
|
||
if self.payload is None:
|
||
return None, None, None
|
||
if persist_call is not None:
|
||
# 模拟 muse_llm.chat 的落库事件形状(window_key 等治理字段此处省略)。
|
||
persist_call(
|
||
{
|
||
"run_id": run_id,
|
||
"caller": caller or "",
|
||
"requested_model_id": FIXED_OPUS_POLICY_ALIAS,
|
||
"actual_model_id": ACTUAL_MODEL,
|
||
"usage": USAGE,
|
||
"prompt": prompt,
|
||
"response": self.payload,
|
||
"role": caller,
|
||
}
|
||
)
|
||
return self.payload, USAGE, ACTUAL_MODEL
|
||
|
||
|
||
def _profile(*, timeout_seconds: float = 30):
|
||
"""构造绑定 WriterDraft v2 schema 的测试 profile(不含任何 CLI 字段)。"""
|
||
|
||
return build_writer_execution_profile(
|
||
max_budget_usd_per_call=Decimal("1.500000"),
|
||
timeout_seconds=timeout_seconds,
|
||
max_context_chars=200000,
|
||
system_prompt="只返回 WriterDraft v2:一个只含 candidateBody 的对象。",
|
||
)
|
||
|
||
|
||
def _success_chat(output: dict | None = None, context: dict | None = None) -> FakeChat:
|
||
"""返回模拟 writer 成功输出的假治理调用。"""
|
||
|
||
payload = output if output is not None else _writer_draft(context or {})
|
||
return FakeChat(json.dumps(payload, ensure_ascii=False))
|
||
|
||
|
||
class RunWriterTest(unittest.TestCase):
|
||
def test_dispatch_contract_injects_prompt_and_frozen_input(self):
|
||
"""派发合同:身份提示与中心角色合同进系统提示词,冻结创作输入作 prompt,支架字段不外泄。"""
|
||
|
||
context = _bound_context()
|
||
chat = _success_chat(context=context)
|
||
profile = _profile()
|
||
result = run_writer(
|
||
context,
|
||
profile=profile,
|
||
governed_chat=chat,
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
|
||
self.assertEqual(result["candidateVersion"], 1)
|
||
self.assertEqual(result["schemaVersion"], "candidate-envelope-v2")
|
||
call = chat.calls[0]
|
||
# 已装配的身份提示与角色合同必须原样注入系统提示词,不得裁剪改写。
|
||
self.assertEqual(call["system"], build_dispatch_system_prompt(profile))
|
||
creative_input = build_writer_creative_input(context)
|
||
self.assertEqual(call["prompt"], canonical_json(creative_input))
|
||
# 创作输入只含创作材料;运行身份与支架字段不得进入模型上下文。
|
||
serialized_input = call["prompt"]
|
||
for forbidden in (
|
||
"runId", "authorizationSnapshot", "retrievalManifest", "contextSnapshot",
|
||
"candidateVersion", "acceptanceEligible", "contentSha256",
|
||
):
|
||
self.assertNotIn(forbidden, serialized_input)
|
||
|
||
def test_failure_modes_fail_closed(self):
|
||
"""链耗尽、底座异常、非法 JSON、schema 违规一律失败关闭并带稳定码。"""
|
||
|
||
context = _bound_context()
|
||
invalid_schema = _writer_draft(context)
|
||
invalid_schema["candidateSha256"] = "sha256:" + "a" * 64
|
||
cases = {
|
||
"WRITER_MODEL_UNAVAILABLE": FakeChat(None),
|
||
"WRITER_RUNTIME_FAILED": FakeChat(RuntimeError("拒绝执行")),
|
||
"WRITER_SCHEMA_INVALID": FakeChat("not-json"),
|
||
"WRITER_SCHEMA_INVALID_EXTRA_FIELD": _success_chat(invalid_schema, context),
|
||
}
|
||
|
||
for expected_code, chat in cases.items():
|
||
code = expected_code.split("_EXTRA")[0]
|
||
with self.subTest(expected_code=expected_code), self.assertRaises(WriterAdapterError) as raised:
|
||
run_writer(
|
||
context,
|
||
profile=_profile(timeout_seconds=1),
|
||
governed_chat=chat,
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
self.assertEqual(raised.exception.code, code)
|
||
self.assertFalse(raised.exception.acceptance_eligible)
|
||
# 底座错误详情不得泄漏进 adapter 错误。
|
||
self.assertNotIn("拒绝执行", str(raised.exception.details))
|
||
|
||
def test_candidate_length_must_fit_dynamic_context_range(self):
|
||
context = _bound_context()
|
||
too_short = _writer_draft(context, body="短" * 3599)
|
||
|
||
with self.assertRaises(WriterAdapterError) as raised:
|
||
run_writer(
|
||
context,
|
||
profile=_profile(),
|
||
governed_chat=_success_chat(too_short, context),
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
|
||
self.assertEqual(raised.exception.code, "candidate_length_out_of_range")
|
||
self.assertEqual(raised.exception.details["actualHanChars"], 3599)
|
||
|
||
def test_production_prompt_range_is_stricter_than_acceptance_floor(self):
|
||
context = _bound_context()
|
||
dynamic = calculate_dynamic_output_contract(
|
||
fine_outline={"targetChars": 4000}, recent_chapter_bodies=[]
|
||
)
|
||
acceptance, generation = build_production_length_contracts(dynamic)
|
||
self.assertEqual(generation["minChars"], 4000)
|
||
self.assertEqual(generation["targetChars"], 7000)
|
||
self.assertEqual(generation["maxChars"], 7000)
|
||
self.assertEqual(acceptance["minChars"], 3001)
|
||
context["outputContract"] = acceptance
|
||
context["generationLengthContract"] = generation
|
||
context["contextSnapshot"]["contextSha256"] = retrieval_identity(context)
|
||
|
||
accepted = run_writer(
|
||
context,
|
||
profile=_profile(),
|
||
governed_chat=_success_chat(_writer_draft(context, body="文" * 3001), context),
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
self.assertEqual(len(accepted["candidateBody"]), 3001)
|
||
self.assertEqual(build_writer_creative_input(context)["lengthContract"], generation)
|
||
|
||
with self.assertRaises(WriterAdapterError) as raised:
|
||
run_writer(
|
||
context,
|
||
profile=_profile(),
|
||
governed_chat=_success_chat(_writer_draft(context, body="文" * 3000), context),
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
self.assertEqual(raised.exception.code, "candidate_length_out_of_range")
|
||
self.assertEqual(raised.exception.details["minChars"], 3001)
|
||
|
||
def test_model_cannot_supply_candidate_hash_or_other_envelope_fields(self):
|
||
context = _bound_context()
|
||
output = _writer_draft(context)
|
||
output["candidateSha256"] = "sha256:" + "a" * 64
|
||
|
||
with self.assertRaises(WriterAdapterError) as raised:
|
||
run_writer(
|
||
context,
|
||
profile=_profile(),
|
||
governed_chat=_success_chat(output, context),
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
|
||
self.assertEqual(raised.exception.code, "WRITER_SCHEMA_INVALID")
|
||
|
||
def test_dynamic_output_contract_uses_outline_density_and_frozen_history(self):
|
||
contract = calculate_dynamic_output_contract(
|
||
fine_outline={
|
||
"hardConstraints": ["事件一", "事件二", "事件三"],
|
||
"foreshadowingActions": [],
|
||
"requiredScenes": [],
|
||
},
|
||
recent_chapter_bodies=["文" * 2501, "文" * 2502, "文" * 2503, "文" * 2504],
|
||
)
|
||
|
||
self.assertEqual(contract["targetChars"], 2100)
|
||
self.assertEqual(contract["minChars"], 2000)
|
||
self.assertEqual(contract["maxChars"], 2730)
|
||
self.assertFalse(contract["frontmatterRequired"])
|
||
self.assertNotIn("newSettingDeclarationRequired", contract)
|
||
|
||
def test_candidate_envelope_metadata_is_bound_by_adapter(self):
|
||
context = _bound_context()
|
||
result = run_writer(
|
||
copy.deepcopy(context),
|
||
profile=_profile(),
|
||
governed_chat=_success_chat(context=context),
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
|
||
self.assertEqual(result["runId"], context["runId"])
|
||
self.assertEqual(result["attempt"], context["attempt"])
|
||
self.assertEqual(result["contextSnapshotId"], context["contextSnapshot"]["manifestId"])
|
||
self.assertEqual(result["contextSnapshotSha256"], context["contextSnapshot"]["contextSha256"])
|
||
self.assertEqual(result["acceptanceEligible"], context["acceptanceEligible"])
|
||
self.assertEqual(
|
||
result["candidateSha256"],
|
||
"sha256:" + hashlib.sha256(result["candidateBody"].encode("utf-8")).hexdigest(),
|
||
)
|
||
|
||
def test_writer_profile_binds_v2_schema_and_prompt_ids(self):
|
||
profile = _profile()
|
||
self.assertEqual(profile.profile_version, "role-writer-v4")
|
||
self.assertEqual(profile.json_schema_id, "writer-draft-v2")
|
||
self.assertEqual(profile.system_prompt_id, "writer-system-prompt-v2")
|
||
self.assertEqual(profile.model_alias, FIXED_OPUS_POLICY_ALIAS)
|
||
self.assertEqual(profile.resolved_model_id, FIXED_OPUS_MODEL_ID)
|
||
self.assertEqual(profile.json_schema["required"], ["candidateBody"])
|
||
self.assertFalse(profile.json_schema["additionalProperties"])
|
||
self.assertTrue(profile.system_prompt.startswith(WRITER_ROLE_PROMPT.rstrip()))
|
||
self.assertIn(ROLE_TASK_SEPARATOR, profile.system_prompt)
|
||
|
||
def test_real_default_requires_explicit_frozen_profile(self):
|
||
"""未显式传入冻结 profile 时,真实默认入口必须失败关闭。"""
|
||
|
||
with self.assertRaises(WriterAdapterError) as raised:
|
||
run_writer(_bound_context())
|
||
|
||
self.assertEqual(raised.exception.code, "WRITER_PROFILE_REQUIRED")
|
||
|
||
def test_run_identity_and_persistence_forward_to_runtime(self):
|
||
"""生产记账:运行号默认随上下文,调用方与存证入口显式转发给治理底座。"""
|
||
|
||
context = _bound_context()
|
||
chat = _success_chat(context=context)
|
||
events = []
|
||
result = run_writer(
|
||
context,
|
||
profile=_profile(),
|
||
governed_chat=chat,
|
||
binding_verifier=lambda _profile: None,
|
||
persist_call=events.append,
|
||
)
|
||
self.assertIn("candidateBody", result)
|
||
self.assertEqual(len(events), 1)
|
||
event = events[0]
|
||
self.assertEqual(event["run_id"], context["runId"])
|
||
self.assertEqual(event["caller"], "writer")
|
||
self.assertEqual(event["role"], "writer")
|
||
self.assertEqual(event["actual_model_id"], ACTUAL_MODEL)
|
||
# prompt 是冻结创作输入的规范 JSON;response 携带模型原始输出。
|
||
self.assertIn("fineOutline", event["prompt"])
|
||
self.assertIn("candidateBody", event["response"])
|
||
|
||
def test_model_cannot_smuggle_draft_through_noncontract_field(self):
|
||
"""模型把草稿塞进非合同字段(如 result)时,缺必需字段必须拒绝。"""
|
||
|
||
context = _bound_context()
|
||
smuggled = {"result": json.dumps(_writer_draft(context), ensure_ascii=False)}
|
||
|
||
with self.assertRaises(WriterAdapterError) as raised:
|
||
run_writer(
|
||
context,
|
||
profile=_profile(),
|
||
governed_chat=FakeChat(json.dumps(smuggled, ensure_ascii=False)),
|
||
binding_verifier=lambda _profile: None,
|
||
)
|
||
|
||
self.assertEqual(raised.exception.code, "WRITER_SCHEMA_INVALID")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|