587 lines
22 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
"""正文智能体 Gate A/B 的确定性唯一判定器。
判定器只消费逐样本脱敏结果,并在可信边界内计算全部统计值。顶层传入的
样本数、平均分、覆盖率或混淆项等聚合字段一律不参与裁决,避免上游通过
伪造汇总改变终态。命中较高优先级终态后不再使用较低优先级条件改写结果。
"""
from __future__ import annotations
import argparse
import json
import math
import sys
from pathlib import Path
from typing import Any, Mapping, Sequence
from writer_rubric import DIMENSIONS
REQUIRED_SCENARIOS = frozenset(
{
"battle",
"character_dialogue",
"turning_point",
"information_reveal",
"returning_character",
}
)
CONFOUND_LABELS = {
"falseNegatives": "假阴",
"falsePositives": "假阳",
"leakage": "泄露",
"reviewerInstability": "评委不稳定",
"newCharactersWithoutCards": "新角色无卡",
}
STABLE_REVIEW_STATUSES = frozenset({"stable_report", "adjudicated_report"})
REVIEW_STATUSES = frozenset({*STABLE_REVIEW_STATUSES, "invalid_unstable"})
def _require_mapping(value: object, field: str) -> Mapping[str, Any]:
"""取得必需对象字段,避免缺失数据被静默当成空对象。"""
if not isinstance(value, Mapping):
raise ValueError(f"{field} 必须是对象")
return value
def _require_samples(value: object) -> list[Mapping[str, Any]]:
"""取得非空逐样本数组,Gate 不接受只有聚合值的输入。"""
if not isinstance(value, list) or not value:
raise ValueError("samples 必须是非空逐样本数组")
if any(not isinstance(item, Mapping) for item in value):
raise ValueError("samples 每项必须是对象")
return list(value)
def _require_integer(value: object, field: str) -> int:
"""取得非负整数计数。"""
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValueError(f"{field} 必须是非负整数")
return value
def _require_number(value: object, field: str) -> float:
"""取得有限数值,显式拒绝布尔值、NaN 和无穷值。"""
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"{field} 必须是有限数值")
result = float(value)
if not math.isfinite(result):
raise ValueError(f"{field} 必须是有限数值")
return result
def _require_ratio(value: object, field: str) -> float:
"""取得闭区间 0-1 内的比例。"""
result = _require_number(value, field)
if not 0.0 <= result <= 1.0:
raise ValueError(f"{field} 必须在 0-1 之间")
return result
def _require_boolean(value: object, field: str) -> bool:
"""取得严格布尔值,拒绝 0/1 等隐式真假值。"""
if not isinstance(value, bool):
raise ValueError(f"{field} 必须是布尔值")
return value
def _sample_identity(sample: Mapping[str, Any], index: int) -> tuple[str, str, str]:
"""校验并返回样本 ID、作品 ID 和场景分类。"""
sample_id = sample.get("sampleId")
if not isinstance(sample_id, str) or not sample_id.strip():
raise ValueError(f"samples[{index}].sampleId 必须是非空字符串")
work_id = sample.get("workId")
if isinstance(work_id, bool) or not isinstance(work_id, (int, str)) or not str(work_id).strip():
raise ValueError(f"samples[{index}].workId 必须是非空整数或字符串")
scenario = sample.get("scenario")
if scenario not in REQUIRED_SCENARIOS:
raise ValueError(f"samples[{index}].scenario 未登记")
return sample_id, str(work_id), str(scenario)
def _sample_status(sample: Mapping[str, Any], index: int) -> dict[str, Any]:
"""从单样本机械结果推导有效性,不接受上游 valid 标记。"""
schema_valid = _require_boolean(sample.get("schemaValid"), f"samples[{index}].schemaValid")
future_leakage = _require_boolean(
sample.get("futureLeakage"), f"samples[{index}].futureLeakage"
)
system_failure = _require_boolean(
sample.get("systemFailure"), f"samples[{index}].systemFailure"
)
review_status = sample.get("reviewStatus")
if review_status not in REVIEW_STATUSES:
raise ValueError(
f"samples[{index}].reviewStatus 必须是 stable_report、"
"adjudicated_report 或 invalid_unstable"
)
return {
"schemaValid": schema_valid,
"futureLeakage": future_leakage,
"systemFailure": system_failure,
"reviewStatus": review_status,
"valid": (
schema_valid
and not future_leakage
and not system_failure
and review_status in STABLE_REVIEW_STATUSES
),
}
def _sample_c_arm(sample: Mapping[str, Any], index: int) -> dict[str, Any]:
"""读取单样本 C 臂机械门结果,聚合时不允许平均掩盖硬错误。"""
c_arm = _require_mapping(sample.get("cArm"), f"samples[{index}].cArm")
return {
"hardConstraintCoverage": _require_ratio(
c_arm.get("hardConstraintCoverage"),
f"samples[{index}].cArm.hardConstraintCoverage",
),
"highSeverityResidualCount": _require_integer(
c_arm.get("highSeverityResidualCount"),
f"samples[{index}].cArm.highSeverityResidualCount",
),
}
def _sample_scores(sample: Mapping[str, Any], index: int) -> dict[str, dict[str, float]]:
"""读取稳定样本的 A/B/C 三臂五维脱敏终分。"""
scores = _require_mapping(sample.get("scores"), f"samples[{index}].scores")
if set(scores) != {"A", "B", "C"}:
raise ValueError(f"samples[{index}].scores 必须精确包含 A/B/C")
result: dict[str, dict[str, float]] = {}
for arm in ("A", "B", "C"):
arm_scores = _require_mapping(scores[arm], f"samples[{index}].scores.{arm}")
if set(arm_scores) != set(DIMENSIONS):
raise ValueError(f"samples[{index}].scores.{arm} 必须精确包含正文五维")
result[arm] = {
dimension: _require_number(
arm_scores[dimension], f"samples[{index}].scores.{arm}.{dimension}"
)
for dimension in DIMENSIONS
}
return result
def _sample_confounds(sample: Mapping[str, Any], index: int) -> dict[str, list[Any]]:
"""校验单样本五类混淆项;顶层聚合混淆项不在信任边界内。"""
source = sample.get("confounders", {})
source = _require_mapping(source, f"samples[{index}].confounders")
unknown = sorted(set(source) - set(CONFOUND_LABELS))
if unknown:
raise ValueError(
f"samples[{index}].confounders 存在未知分类: {','.join(unknown)}"
)
result: dict[str, list[Any]] = {}
for category in CONFOUND_LABELS:
items = source.get(category, [])
if not isinstance(items, list):
raise ValueError(f"samples[{index}].confounders.{category} 必须是数组")
if any(not isinstance(item, (str, Mapping)) for item in items):
raise ValueError(
f"samples[{index}].confounders.{category} 只能包含字符串或对象"
)
result[category] = list(items)
return result
def _aggregate_confounds(samples: Sequence[Mapping[str, Any]]) -> dict[str, list[Any]]:
"""逐样本汇总五类混淆项,并补入可机械推导的泄露与不稳定记录。"""
aggregated = {category: [] for category in CONFOUND_LABELS}
for index, sample in enumerate(samples):
sample_id, _, _ = _sample_identity(sample, index)
status = _sample_status(sample, index)
for category, items in _sample_confounds(sample, index).items():
for item in items:
if isinstance(item, str):
aggregated[category].append({"sampleId": sample_id, "detail": item})
else:
aggregated[category].append({**dict(item), "sampleId": sample_id})
if status["futureLeakage"] and not any(
item.get("sampleId") == sample_id for item in aggregated["leakage"]
):
aggregated["leakage"].append(
{"sampleId": sample_id, "code": "future_leakage_detected"}
)
if status["reviewStatus"] == "invalid_unstable" and not any(
item.get("sampleId") == sample_id
for item in aggregated["reviewerInstability"]
):
aggregated["reviewerInstability"].append(
{"sampleId": sample_id, "code": "invalid_unstable"}
)
return aggregated
def _report(
gate: str,
status: str,
reasons: list[str],
confounds: dict[str, list[Any]],
metrics: Mapping[str, Any],
) -> dict[str, Any]:
"""组装稳定报告结构;reasons 顺序同时表达当前层内优先级。"""
return {
"schemaVersion": "writer-gate-report-v1",
"gate": gate,
"status": status,
"primaryReason": reasons[0],
"reasons": reasons,
"metrics": dict(metrics),
"confounds": confounds,
}
def _gate_a_metrics(samples: Sequence[Mapping[str, Any]]) -> dict[str, Any]:
"""完全从逐样本结果计算 Gate A 统计。"""
statuses: list[dict[str, Any]] = []
c_arms: list[dict[str, Any]] = []
seen_ids: set[str] = set()
for index, sample in enumerate(samples):
sample_id, _, _ = _sample_identity(sample, index)
if sample_id in seen_ids:
raise ValueError(f"sampleId 重复: {sample_id}")
seen_ids.add(sample_id)
status = _sample_status(sample, index)
statuses.append(status)
# schema、泄漏或系统失败可能发生在 C 臂生成前,此时没有 cArm 是合法的
# 逐样本失败形态,不能让缺失的下游结果覆盖 Gate 的既定裁决顺序。
if (
status["schemaValid"]
and not status["futureLeakage"]
and not status["systemFailure"]
):
c_arms.append(_sample_c_arm(sample, index))
return {
"sampleCount": len(samples),
"validSampleCount": sum(item["valid"] for item in statuses),
"invalidSchemaCount": sum(not item["schemaValid"] for item in statuses),
"futureLeakageCount": sum(item["futureLeakage"] for item in statuses),
"systemFailureCount": sum(item["systemFailure"] for item in statuses),
"unstableSampleCount": sum(
item["reviewStatus"] == "invalid_unstable" for item in statuses
),
"cArmHighSeverityResidualCount": sum(
item["highSeverityResidualCount"] for item in c_arms
),
"cArmHardConstraintCoverage": min(
(item["hardConstraintCoverage"] for item in c_arms), default=1.0
),
}
def _short_circuit_confounds(gate_input: Mapping[str, Any]) -> dict[str, list[Any]]:
"""Gate A 已决定 Gate B 终态时,尽量保留合法逐样本混淆项。"""
if "samples" not in gate_input:
return {category: [] for category in CONFOUND_LABELS}
try:
return _aggregate_confounds(_require_samples(gate_input.get("samples")))
except ValueError:
# 混淆项属于报告信息,不能反向推翻更高优先级的 Gate A 终态。
return {category: [] for category in CONFOUND_LABELS}
def decide_gate_a(gate_input: Mapping[str, Any]) -> dict[str, Any]:
"""按“证据不足 -> 硬失败 -> 通过”唯一顺序裁决 Gate A。"""
gate_input = _require_mapping(gate_input, "Gate A 输入")
samples = _require_samples(gate_input.get("samples"))
normalized = _gate_a_metrics(samples)
confounds = _aggregate_confounds(samples)
if normalized["validSampleCount"] < 5:
return _report(
"A",
"insufficient_evidence",
["valid_sample_count_below_5"],
confounds,
normalized,
)
failure_reasons: list[str] = []
if normalized["invalidSchemaCount"] > 0:
failure_reasons.append("schema_invalid")
if normalized["futureLeakageCount"] > 0:
failure_reasons.append("future_leakage")
if normalized["systemFailureCount"] > 0:
failure_reasons.append("system_failure")
if normalized["cArmHighSeverityResidualCount"] > 0:
failure_reasons.append("c_arm_high_severity_residual")
if normalized["cArmHardConstraintCoverage"] < 1.0:
failure_reasons.append("c_arm_hard_constraint_coverage_below_100_percent")
if failure_reasons:
return _report("A", "failed", failure_reasons, confounds, normalized)
return _report(
"A", "passed", ["all_gate_a_conditions_met"], confounds, normalized
)
def _gate_b_metrics(samples: Sequence[Mapping[str, Any]]) -> dict[str, Any]:
"""完全从逐样本结果计算 Gate B 充分性、退化和增益指标。"""
work_counts: dict[str, int] = {}
scenarios: set[str] = set()
statuses: list[dict[str, Any]] = []
c_arms: list[dict[str, Any]] = []
deltas: dict[str, list[dict[str, float]]] = {"B-A": [], "C-A": []}
seen_ids: set[str] = set()
for index, sample in enumerate(samples):
sample_id, work_id, scenario = _sample_identity(sample, index)
if sample_id in seen_ids:
raise ValueError(f"sampleId 重复: {sample_id}")
seen_ids.add(sample_id)
work_counts[work_id] = work_counts.get(work_id, 0) + 1
scenarios.add(scenario)
status = _sample_status(sample, index)
statuses.append(status)
c_arms.append(_sample_c_arm(sample, index))
if status["reviewStatus"] in STABLE_REVIEW_STATUSES:
scores = _sample_scores(sample, index)
for comparison, arm in (("B-A", "B"), ("C-A", "C")):
deltas[comparison].append(
{
dimension: scores[arm][dimension] - scores["A"][dimension]
for dimension in DIMENSIONS
}
)
total_count = len(samples)
stable_count = len(deltas["C-A"])
averages = {
comparison: {
dimension: (
sum(item[dimension] for item in comparison_deltas) / stable_count
if stable_count
else 0.0
)
for dimension in DIMENSIONS
}
for comparison, comparison_deltas in deltas.items()
}
# Gate B 的正式退化与增益阈值只消费 C-A;B-A 仅作为卡索引单线的诊断信息。
c_minus_a_deltas = deltas["C-A"]
decline_ratios = {
dimension: (
sum(item[dimension] < -0.5 for item in c_minus_a_deltas) / stable_count
if stable_count
else 0.0
)
for dimension in DIMENSIONS
}
any_decline_ratio = (
sum(
any(item[dimension] < -0.5 for dimension in DIMENSIONS)
for item in c_minus_a_deltas
)
/ stable_count
if stable_count
else 0.0
)
fidelity_positive_ratio = (
sum(
item["setting_entity_fidelity"] > 0 for item in c_minus_a_deltas
)
/ stable_count
if stable_count
else 0.0
)
return {
"workSampleCounts": dict(sorted(work_counts.items())),
"workCount": len(work_counts),
"totalSampleCount": total_count,
"stableSampleCount": stable_count,
"coveredScenarios": sorted(scenarios),
"unstableSampleRatio": sum(
item["reviewStatus"] == "invalid_unstable" for item in statuses
)
/ total_count,
"cArmHardConstraintCoverage": min(
(item["hardConstraintCoverage"] for item in c_arms), default=0.0
),
"cArmHighSeverityResidualCount": sum(
item["highSeverityResidualCount"] for item in c_arms
),
"averageDeltas": averages,
"bMinusAAverageDeltas": averages["B-A"],
"cMinusAAverageDeltas": averages["C-A"],
"dimensionDeclineOverHalfRatios": decline_ratios,
"anyDimensionDeclineOverHalfRatio": any_decline_ratio,
"fidelityPositiveSampleRatio": fidelity_positive_ratio,
}
def decide_gate_b(gate_input: Mapping[str, Any]) -> dict[str, Any]:
"""按 Gate A、证据、退化、增益四层唯一顺序裁决 Gate B。"""
gate_input = _require_mapping(gate_input, "Gate B 输入")
samples = _require_samples(gate_input.get("samples"))
# 顶层 gateAStatus 不在可信边界内;必须用 Gate B 的同批逐样本结果重算。
gate_a_report = decide_gate_a({"gate": "A", "samples": samples})
gate_a_status = gate_a_report["status"]
gate_a_metrics = {
"recomputedGateAStatus": gate_a_status,
"recomputedGateAMetrics": gate_a_report["metrics"],
}
if gate_a_status == "insufficient_evidence":
return _report(
"B",
"insufficient_evidence",
["gate_a_insufficient_evidence"],
gate_a_report["confounds"],
gate_a_metrics,
)
if gate_a_status == "failed":
return _report(
"B",
"failed",
["gate_a_failed"],
gate_a_report["confounds"],
gate_a_metrics,
)
normalized = _gate_b_metrics(samples)
normalized.update(gate_a_metrics)
confounds = _aggregate_confounds(samples)
evidence_reasons: list[str] = []
if normalized["workCount"] < 2:
evidence_reasons.append("work_count_below_2")
if any(count < 5 for count in normalized["workSampleCounts"].values()):
evidence_reasons.append("per_work_sample_count_below_5")
if normalized["totalSampleCount"] < 10:
evidence_reasons.append("total_sample_count_below_10")
if not REQUIRED_SCENARIOS.issubset(normalized["coveredScenarios"]):
evidence_reasons.append("scenario_coverage_incomplete")
if normalized["unstableSampleRatio"] > 0.20:
evidence_reasons.append("unstable_sample_ratio_above_20_percent")
if evidence_reasons:
return _report(
"B", "insufficient_evidence", evidence_reasons, confounds, normalized
)
quality_reasons: list[str] = []
if normalized["cArmHardConstraintCoverage"] < 1.0:
quality_reasons.append("c_arm_hard_constraint_coverage_below_100_percent")
if normalized["cArmHighSeverityResidualCount"] > 0:
quality_reasons.append("c_arm_high_severity_residual")
averages = normalized["averageDeltas"]["C-A"]
if averages["style_consistency"] < -0.25:
quality_reasons.append("style_consistency_average_delta_below_minus_0_25")
if averages["narrative_tension"] < -0.25:
quality_reasons.append("narrative_tension_average_delta_below_minus_0_25")
if normalized["anyDimensionDeclineOverHalfRatio"] > 0.20:
quality_reasons.append("dimension_decline_over_half_ratio_above_20_percent")
if quality_reasons:
return _report("B", "failed", quality_reasons, confounds, normalized)
if (
averages["setting_entity_fidelity"] >= 0.25
and normalized["fidelityPositiveSampleRatio"] >= 0.60
):
return _report(
"B", "passed", ["fidelity_gain_threshold_met"], confounds, normalized
)
return _report(
"B", "no_gain", ["fidelity_gain_threshold_not_met"], confounds, normalized
)
def decide_gate(gate_input: Mapping[str, Any]) -> dict[str, Any]:
"""从带 gate 字段的统一输入分派 Gate A 或 Gate B。"""
gate_input = _require_mapping(gate_input, "gate-input.json")
gate = gate_input.get("gate")
if gate == "A":
return decide_gate_a(gate_input)
if gate == "B":
return decide_gate_b(gate_input)
raise ValueError("gate 必须是 A 或 B")
def render_summary(report: Mapping[str, Any]) -> str:
"""渲染只含内部统计、终态和混淆项的 Markdown 摘要。"""
lines = [
f"# Writer Gate {report['gate']} 裁决摘要",
"",
f"- 终态:`{report['status']}`",
f"- 主原因:`{report['primaryReason']}`",
f"- 原因码:`{', '.join(report['reasons'])}`",
"",
"## 混淆项",
"",
]
confounds = _require_mapping(report.get("confounds"), "report.confounds")
for category, label in CONFOUND_LABELS.items():
lines.append(f"### {label}")
lines.append("")
items = confounds.get(category, [])
if not items:
lines.append("- 未观察到或未报告")
else:
for item in items:
rendered = (
item
if isinstance(item, str)
else json.dumps(item, ensure_ascii=False, sort_keys=True)
)
lines.append(f"- {rendered}")
lines.append("")
return "\n".join(lines)
def _parse_args() -> argparse.Namespace:
"""解析任务计划约定的运行目录与摘要输出参数。"""
parser = argparse.ArgumentParser(description="裁决正文智能体 Gate A/B")
parser.add_argument(
"--run-dir", type=Path, required=True, help="包含 gate-input.json 的回放运行目录"
)
parser.add_argument("--summary-output", type=Path, help="可选的脱敏 Markdown 摘要输出路径")
return parser.parse_args()
def main() -> int:
"""执行离线判定,并写出唯一 JSON 终态和可选摘要。"""
args = _parse_args()
input_path = args.run_dir / "gate-input.json"
output_path = args.run_dir / "gate-report.json"
try:
gate_input = json.loads(input_path.read_text(encoding="utf-8"))
report = decide_gate(gate_input)
output_path.write_text(
json.dumps(report, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
if args.summary_output is not None:
args.summary_output.parent.mkdir(parents=True, exist_ok=True)
args.summary_output.write_text(render_summary(report), encoding="utf-8")
except (OSError, ValueError, json.JSONDecodeError) as error:
print(f"writer_gate: {error}", file=sys.stderr)
return 2
return 0
if __name__ == "__main__":
raise SystemExit(main())