713 lines
27 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
"""正文候选的有限补证、重写、机械审查与 CAS 编排。"""
from __future__ import annotations
import copy
import os
import pathlib
import sys
import tempfile
from dataclasses import dataclass
from typing import Any, Callable, Mapping, Protocol, Sequence
SCRIPT_DIR = pathlib.Path(__file__).resolve().parent
SKILLS_DIR = SCRIPT_DIR.parents[1]
DETECT_DIR = SKILLS_DIR / "detect" / "scripts"
READ_CONTEXT_DIR = SKILLS_DIR / "read-context" / "scripts"
for path in (DETECT_DIR, READ_CONTEXT_DIR):
if str(path) not in sys.path:
sys.path.insert(0, str(path))
from check_writer_candidate import check_writer_candidate # noqa: E402
from writer_contract import ( # noqa: E402
ContractError,
canonical_json,
retrieval_identity,
validate_writer_context,
validate_writer_output,
)
MAX_EVIDENCE_REQUESTS = 3
MAX_REWRITES = 2
class PipelineError(RuntimeError):
"""携带稳定失败码和终态结果的编排错误。"""
def __init__(
self,
code: str,
message: str,
*,
details: Mapping[str, Any] | None = None,
result: Mapping[str, Any] | None = None,
):
super().__init__(message)
self.code = code
self.details = dict(details or {})
self.result = dict(result or {})
self.acceptance_eligible = False
@dataclass(frozen=True)
class CasToken:
"""一次候选状态的不可变 CAS 身份。"""
run_id: str
attempt: int
candidate_version: int
state: str
revision: int
class CasStateStore(Protocol):
"""生产存储必须实现的最小 CAS 接口;本轮测试只用内存 fake。"""
def create(self, run_id: str, *, attempt: int, candidate_version: int) -> CasToken: ...
def transition(self, expected: CasToken, target_state: str) -> CasToken | None: ...
def start_next(
self, expected: CasToken, *, attempt: int, candidate_version: int
) -> CasToken | None: ...
def latest(self, run_id: str) -> CasToken | None: ...
class SemanticDetector(Protocol):
"""语义 detector 的明确边界;实现方不得读写状态或候选文件。"""
def __call__(
self,
context: Mapping[str, Any],
candidate: Mapping[str, Any],
mechanical_report: Mapping[str, Any],
) -> Mapping[str, Any]: ...
class InMemoryCasStateStore:
"""仅供实验和测试使用的线程外内存 CAS fake。"""
_ALLOWED = {
"DRAFT": frozenset({"CHECKING"}),
"CHECKING": frozenset({"PASSED", "REJECTED"}),
"PASSED": frozenset(),
"REJECTED": frozenset(),
}
def __init__(self) -> None:
self._latest: dict[str, CasToken] = {}
def create(self, run_id: str, *, attempt: int, candidate_version: int) -> CasToken:
"""只允许为未存在的 run 创建第一条 DRAFT。"""
if run_id in self._latest:
raise PipelineError("CAS_CONFLICT", "run 已存在,不能重复创建初始状态")
token = CasToken(run_id, attempt, candidate_version, "DRAFT", 1)
self._latest[run_id] = token
return token
def transition(self, expected: CasToken, target_state: str) -> CasToken | None:
"""仅当最新 token 完全匹配且转换合法时更新。"""
current = self._latest.get(expected.run_id)
if current != expected or target_state not in self._ALLOWED.get(expected.state, frozenset()):
return None
updated = CasToken(
expected.run_id,
expected.attempt,
expected.candidate_version,
target_state,
expected.revision + 1,
)
self._latest[expected.run_id] = updated
return updated
def start_next(
self, expected: CasToken, *, attempt: int, candidate_version: int
) -> CasToken | None:
"""只允许从最新 REJECTED 创建单调递增的新 DRAFT。"""
current = self._latest.get(expected.run_id)
if (
current != expected
or expected.state != "REJECTED"
or attempt <= expected.attempt
or candidate_version <= expected.candidate_version
):
return None
updated = CasToken(
expected.run_id,
attempt,
candidate_version,
"DRAFT",
expected.revision + 1,
)
self._latest[expected.run_id] = updated
return updated
def latest(self, run_id: str) -> CasToken | None:
"""返回 run 的最新不可变状态。"""
return self._latest.get(run_id)
def atomic_write_json(path: str | pathlib.Path, value: Mapping[str, Any]) -> None:
"""先写同目录临时文件并 fsync,再以原子替换发布正式结果。"""
target = pathlib.Path(path)
target.parent.mkdir(parents=True, exist_ok=True)
descriptor, temporary_name = tempfile.mkstemp(
prefix=f".{target.name}.", suffix=".tmp", dir=target.parent
)
temporary = pathlib.Path(temporary_name)
try:
with os.fdopen(descriptor, "w", encoding="utf-8", newline="\n") as handle:
handle.write(canonical_json(value))
handle.write("\n")
handle.flush()
os.fsync(handle.fileno())
os.replace(str(temporary), str(target))
except BaseException:
# replace 成功后临时路径已不存在;失败时只清理本次临时文件。
try:
temporary.unlink(missing_ok=True)
finally:
raise
def _cas_or_fail(token: CasToken | None, action: str) -> CasToken:
"""把所有 CAS 竞争统一为稳定失败码。"""
if token is None:
raise PipelineError("CAS_CONFLICT", f"状态 CAS 失败: {action}")
return token
def _next_context_without_new_evidence(
context: Mapping[str, Any], *, attempt: int
) -> dict[str, Any]:
"""为修复重写创建新 attempt,并重新绑定上下文 hash。"""
advanced = copy.deepcopy(dict(context))
advanced["attempt"] = attempt
advanced["contextSnapshot"]["contextSha256"] = retrieval_identity(advanced)
try:
return validate_writer_context(advanced)
except ContractError as exc:
raise PipelineError("REASSEMBLED_CONTEXT_INVALID", str(exc)) from exc
def _validate_semantic_report(
report: Mapping[str, Any], candidate: Mapping[str, Any]
) -> list[dict[str, Any]]:
"""校验 fake/未来模型 detector 的候选绑定和阻塞项结构。"""
if not isinstance(report, Mapping):
raise PipelineError("SEMANTIC_DETECTOR_INVALID", "语义 detector 必须返回对象")
if report.get("status") not in {"passed", "rejected"}:
raise PipelineError("SEMANTIC_DETECTOR_INVALID", "语义 detector status 非法")
if (
report.get("candidateVersion") != candidate["candidateVersion"]
or report.get("candidateSha256") != candidate["candidateSha256"]
):
raise PipelineError("SEMANTIC_DETECTOR_STALE", "语义 detector 结果绑定了旧候选")
failures = report.get("blockingFailures")
if not isinstance(failures, list) or any(not isinstance(item, Mapping) for item in failures):
raise PipelineError("SEMANTIC_DETECTOR_INVALID", "语义 detector blockingFailures 非法")
if any(not isinstance(item.get("code"), str) or not item["code"] for item in failures):
raise PipelineError("SEMANTIC_DETECTOR_INVALID", "语义 detector 阻塞项缺少稳定 code")
if report["status"] == "passed" and failures:
raise PipelineError("SEMANTIC_DETECTOR_INVALID", "语义 detector 通过时不得含阻塞项")
if report["status"] == "rejected" and not failures:
raise PipelineError("SEMANTIC_DETECTOR_INVALID", "语义 detector 拒绝时必须包含阻塞项")
return [dict(item) for item in failures]
def _validate_mechanical_report(
report: Any, candidate: Mapping[str, Any]
) -> list[dict[str, Any]]:
"""严格校验机械 detector 报告版本、候选绑定和通过状态。"""
if not isinstance(report, Mapping):
raise PipelineError("MECHANICAL_DETECTOR_INVALID", "机械 detector 必须返回对象")
if report.get("schemaVersion") != "writer-detector-report-v1":
raise PipelineError("MECHANICAL_DETECTOR_INVALID", "机械 detector 报告版本非法")
expected = {
"runId": candidate.get("runId"),
"attempt": candidate.get("attempt"),
"candidateVersion": candidate.get("candidateVersion"),
"candidateSha256": candidate.get("candidateSha256"),
}
if any(report.get(field) != value for field, value in expected.items()):
raise PipelineError("MECHANICAL_DETECTOR_INVALID", "机械 detector 报告绑定了旧候选")
passed = report.get("passed")
failures = report.get("blockingFailures")
if not isinstance(passed, bool):
raise PipelineError("MECHANICAL_DETECTOR_INVALID", "机械 detector passed 必须是布尔值")
if not isinstance(failures, list) or any(not isinstance(item, Mapping) for item in failures):
raise PipelineError("MECHANICAL_DETECTOR_INVALID", "机械 detector blockingFailures 非法")
if any(not isinstance(item.get("code"), str) or not item["code"] for item in failures):
raise PipelineError("MECHANICAL_DETECTOR_INVALID", "机械 detector 阻塞项缺少稳定 code")
if passed != (not failures):
raise PipelineError("MECHANICAL_DETECTOR_INVALID", "机械 detector 通过状态与阻塞项矛盾")
return [dict(item) for item in failures]
def _terminal_result(
*,
run_id: str,
status: str,
token: CasToken,
evidence_request_count: int,
rewrite_count: int,
trace: Sequence[Mapping[str, Any]],
candidate: Mapping[str, Any] | None,
failure_code: str | None = None,
) -> dict[str, Any]:
"""生成可原子落盘的稳定终态结果。"""
result = {
"schemaVersion": "writer-pipeline-result-v1",
"runId": run_id,
"status": status,
"attempt": token.attempt,
"candidateVersion": token.candidate_version,
"candidateSha256": candidate.get("candidateSha256") if candidate else None,
"failureCode": failure_code,
"evidenceRequestCount": evidence_request_count,
"rewriteCount": rewrite_count,
"trace": [dict(item) for item in trace],
}
# 只有通过全部合同和审查门的终态才暴露可进入 Shadow 的最终候选。
if status == "PASSED" and candidate is not None:
result["candidateArtifact"] = dict(candidate)
return result
def _publish_if_requested(
result_path: str | pathlib.Path | None, result: Mapping[str, Any]
) -> None:
"""仅在调用方显式给出路径时发布终态文件。"""
if result_path is not None:
atomic_write_json(result_path, result)
def _terminal_failure(
*,
code: str,
message: str,
token: CasToken,
state_store: CasStateStore,
run_id: str,
evidence_request_count: int,
rewrite_count: int,
trace: Sequence[Mapping[str, Any]],
candidate: Mapping[str, Any] | None,
result_path: str | pathlib.Path | None,
details: Mapping[str, Any] | None = None,
) -> PipelineError:
"""将已进入检查的失败统一收敛为可发布的 REJECTED 终态。"""
if token.state != "CHECKING":
raise PipelineError("CAS_CONFLICT", f"失败收敛不接受状态: {token.state}")
rejected = _cas_or_fail(
state_store.transition(token, "REJECTED"), "CHECKING -> REJECTED"
)
terminal_trace = [dict(item) for item in trace]
terminal_trace.append(
{
"attempt": rejected.attempt,
"candidateVersion": rejected.candidate_version,
"candidateSha256": candidate.get("candidateSha256") if candidate else None,
"status": "pipeline_failed",
"failureCodes": [code],
}
)
result = _terminal_result(
run_id=run_id,
status="REJECTED",
token=rejected,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=terminal_trace,
candidate=candidate,
failure_code=code,
)
_publish_if_requested(result_path, result)
return PipelineError(code, message, details=details, result=result)
def run_writer_pipeline(
*,
context: Mapping[str, Any],
requirements: Mapping[str, Any],
writer: Callable[[Mapping[str, Any], int, list[dict[str, Any]]], Mapping[str, Any]],
evidence_provider: Callable[
[Mapping[str, Any], list[dict[str, Any]], int], Mapping[str, Any]
],
semantic_detector: SemanticDetector,
state_store: CasStateStore,
result_path: str | pathlib.Path | None = None,
) -> dict[str, Any]:
"""执行最多 3 次补证、2 次重写的正文候选有限收敛循环。"""
try:
current_context = validate_writer_context(context)
except ContractError as exc:
raise PipelineError("WRITER_CONTEXT_INVALID", str(exc)) from exc
run_id = current_context["runId"]
candidate_version = 1
evidence_request_count = 0
rewrite_count = 0
repair_failures: list[dict[str, Any]] = []
trace: list[dict[str, Any]] = []
draft = state_store.create(
run_id,
attempt=current_context["attempt"],
candidate_version=candidate_version,
)
while True:
checking = _cas_or_fail(state_store.transition(draft, "CHECKING"), "DRAFT -> CHECKING")
try:
raw_candidate = writer(current_context, candidate_version, repair_failures)
except PipelineError as exc:
raise _terminal_failure(
code=exc.code,
message=str(exc),
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=None,
result_path=result_path,
details=exc.details,
) from exc
except Exception as exc:
raise _terminal_failure(
code="WRITER_FAILED",
message="fake/adapter 写手调用失败",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=None,
result_path=result_path,
) from exc
if not isinstance(raw_candidate, Mapping):
raise _terminal_failure(
code="WRITER_OUTPUT_INVALID",
message="写手必须返回对象",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=None,
result_path=result_path,
)
candidate = dict(raw_candidate)
if (
candidate.get("runId") != run_id
or candidate.get("attempt") != current_context["attempt"]
or candidate.get("candidateVersion") != candidate_version
):
raise _terminal_failure(
code="WRITER_OUTPUT_STALE",
message="写手返回了旧 attempt 或 candidateVersion",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
)
try:
candidate = validate_writer_output(candidate)
except ContractError as exc:
# 合同错误作为可定位审查失败进入有限重写,而不是绕过状态机。
try:
mechanical_report = check_writer_candidate(current_context, candidate, requirements)
candidate_failures = _validate_mechanical_report(mechanical_report, candidate)
except PipelineError as detector_exc:
raise _terminal_failure(
code=detector_exc.code,
message=str(detector_exc),
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
details=detector_exc.details,
) from detector_exc
except Exception as detector_exc:
raise _terminal_failure(
code="MECHANICAL_DETECTOR_FAILED",
message="机械 detector 调用失败",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
) from detector_exc
if not candidate_failures:
candidate_failures = [{"code": "OUTPUT_CONTRACT_INVALID", "message": str(exc)}]
else:
candidate_failures = []
mechanical_report = {}
requests = candidate.get("evidenceRequests", [])
if not isinstance(requests, list):
requests = []
if evidence_request_count + len(requests) > MAX_EVIDENCE_REQUESTS:
raise _terminal_failure(
code="EVIDENCE_REQUEST_LIMIT_REACHED",
message="补证请求累计超过 3 次",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
)
if requests:
evidence_request_count += len(requests)
if rewrite_count >= MAX_REWRITES:
raise _terminal_failure(
code="REWRITE_LIMIT_REACHED",
message="补证后重写已达到 2 次",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
)
next_attempt = current_context["attempt"] + 1
try:
provided = evidence_provider(current_context, [dict(item) for item in requests], next_attempt)
current_context = validate_writer_context(provided)
except Exception as exc:
raise _terminal_failure(
code="REASSEMBLED_CONTEXT_INVALID",
message=str(exc),
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
) from exc
if current_context["runId"] != run_id or current_context["attempt"] != next_attempt:
raise _terminal_failure(
code="REASSEMBLED_CONTEXT_INVALID",
message="补证上下文 runId/attempt 不匹配",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
)
rejected = _cas_or_fail(
state_store.transition(checking, "REJECTED"), "CHECKING -> REJECTED"
)
rewrite_count += 1
candidate_version += 1
trace.append(
{
"attempt": checking.attempt,
"candidateVersion": checking.candidate_version,
"status": "evidence_requested",
"requestIds": [item.get("requestId") for item in requests],
}
)
repair_failures = []
draft = _cas_or_fail(
state_store.start_next(
rejected,
attempt=next_attempt,
candidate_version=candidate_version,
),
"REJECTED -> next DRAFT",
)
continue
if not mechanical_report:
try:
mechanical_report = check_writer_candidate(current_context, candidate, requirements)
candidate_failures = _validate_mechanical_report(mechanical_report, candidate)
except PipelineError as detector_exc:
raise _terminal_failure(
code=detector_exc.code,
message=str(detector_exc),
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
details=detector_exc.details,
) from detector_exc
except Exception as exc:
raise _terminal_failure(
code="MECHANICAL_DETECTOR_FAILED",
message="机械 detector 调用失败",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
) from exc
semantic_report: Mapping[str, Any] | None = None
if not candidate_failures:
try:
semantic_report = semantic_detector(current_context, candidate, mechanical_report)
candidate_failures.extend(_validate_semantic_report(semantic_report, candidate))
except PipelineError as exc:
raise _terminal_failure(
code=exc.code,
message=str(exc),
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
details=exc.details,
) from exc
except Exception as exc:
raise _terminal_failure(
code="SEMANTIC_DETECTOR_FAILED",
message="fake/语义 detector 调用失败",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
) from exc
trace.append(
{
"attempt": checking.attempt,
"candidateVersion": checking.candidate_version,
"candidateSha256": candidate.get("candidateSha256"),
"mechanicalPassed": mechanical_report.get("passed", False),
"semanticStatus": semantic_report.get("status") if semantic_report else "not_run",
"failureCodes": [item.get("code") for item in candidate_failures],
# 保留接受层审计所需的原始报告;摘要字段仍供稳定绑定检查使用。
"mechanicalReport": dict(mechanical_report),
"semanticReport": dict(semantic_report) if semantic_report else None,
}
)
if not candidate_failures:
passed = _cas_or_fail(
state_store.transition(checking, "PASSED"), "CHECKING -> PASSED"
)
result = _terminal_result(
run_id=run_id,
status="PASSED",
token=passed,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
)
_publish_if_requested(result_path, result)
return result
if rewrite_count >= MAX_REWRITES:
raise _terminal_failure(
code="REWRITE_LIMIT_REACHED",
message="候选两次重写后仍未通过",
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
)
rewrite_count += 1
candidate_version += 1
next_attempt = current_context["attempt"] + 1
try:
next_context = _next_context_without_new_evidence(
current_context, attempt=next_attempt
)
except PipelineError as exc:
raise _terminal_failure(
code=exc.code,
message=str(exc),
token=checking,
state_store=state_store,
run_id=run_id,
evidence_request_count=evidence_request_count,
rewrite_count=rewrite_count,
trace=trace,
candidate=candidate,
result_path=result_path,
details=exc.details,
) from exc
rejected = _cas_or_fail(
state_store.transition(checking, "REJECTED"), "CHECKING -> REJECTED"
)
current_context = next_context
repair_failures = [dict(item) for item in candidate_failures]
draft = _cas_or_fail(
state_store.start_next(
rejected,
attempt=next_attempt,
candidate_version=candidate_version,
),
"REJECTED -> next DRAFT",
)
__all__ = [
"MAX_EVIDENCE_REQUESTS",
"MAX_REWRITES",
"CasToken",
"CasStateStore",
"InMemoryCasStateStore",
"PipelineError",
"SemanticDetector",
"atomic_write_json",
"run_writer_pipeline",
]