diff --git a/docs/reference/configuration.md b/docs/reference/configuration.md index 8c28d829..532bd00a 100644 --- a/docs/reference/configuration.md +++ b/docs/reference/configuration.md @@ -7,7 +7,7 @@ WaveBench 从指定路径或默认的 `wavebench.toml` 加载本地实验台配 | 类别 | 表 | | --- | --- | | 必需 | `[connection]`、`[scope]` | -| 可选 | `[autoscale]`、`[waveform]`、`[output]`、`[quality]`、`[safety_limits]`、`[tui]`、`[source]`、`[rf_source]`、`[power]`、`[dmm]` | +| 可选 | `[autoscale]`、`[waveform]`、`[output]`、`[quality]`、`[safety_limits]`、`[tui]`、`[source]`、`[rf_source]`、`[power]`、`[dmm]`、`[advisor]` | 字段、默认值和跨字段约束由 config model 与 parser 定义。修改 plan 或配置前,先核对[示例配置](https://github.com/Scaxlibur/wavebench/blob/master/wavebench.example.toml)、当前 CLI help 和相关 Reference。 @@ -37,6 +37,14 @@ access = "read_only" 配置中的 `access` 不能替代真实接线、操作系统权限或仪器自身保护。 +## Advisor 外发边界 + +`[advisor]` 控制 advisor(外部判断模型)能否把状态发出本机。默认 `enabled = false`,关闭时不会构造任何外发请求;开启必须同时给出 `endpoint_hosts` 与 `allowed_state_fields` 白名单,否则配置校验直接失败。`accept`/`review` 是「概率 → 动作」的阈值,由 Core 拥有,必须满足 `0 <= review <= accept <= 1`。 + +外发还需通过同意门:发送前产出完整预览(内容、字节数、sha256、逐条不可信来源),预览不联网;同意默认按次,可登记为同一 run 内有效,并绑定 endpoint 集合与允许字段集,任一变化即失效;非交互场景(run plan、CI、MCP)一律拒绝。结论写入对应 run 目录的 `decisions/`(schema `wavebench.decision.v1`),附加式,不影响 `run.json`、质量门或 `auto_recover`。 + +实现状态与分期见 [advisor 插件类别 RFC(PR #20)](https://github.com/Scaxlibur/WaveBench/pull/20)(`Draft`)。 + ## 安全限制 `[safety_limits]` 用于在打开 transport 前限制 Source/Power 写入。Source V2 的端口电压下界和上界必须成对出现,并且下界小于上界。限制不应为了通过一次实验而放宽;应先确认实际端接、量程和实验要求。 diff --git a/src/wavebench/config.py b/src/wavebench/config.py index 9076ac58..3ec0fc8d 100644 --- a/src/wavebench/config.py +++ b/src/wavebench/config.py @@ -8,6 +8,7 @@ from .errors import ConfigError from .services.access_policy import AccessMode, normalize_access_mode +from .services.advisor_consent import normalize_names WAVEFORM_POINTS_ALIASES = { "def": "DEF", @@ -144,6 +145,34 @@ class DmmConfig: dsrdtr: bool = False access: AccessMode = "read_write" +@dataclass(frozen=True) +class AdvisorConfig: + """advisor(外部判断模型)配置。 + + 默认关闭;开启必须同时给出 endpoint 与字段白名单,避免"打开开关即放宽外发范围"。 + 概率阈值由 Core 拥有,插件不得自带。 + """ + + enabled: bool = False + endpoint_hosts: tuple[str, ...] = () + allowed_state_fields: tuple[str, ...] = () + accept: float = 0.60 + review: float = 0.35 + + def __post_init__(self) -> None: + if isinstance(self.accept, bool) or isinstance(self.review, bool): + raise ConfigError("advisor thresholds must be numbers") + if not 0.0 <= float(self.review) <= float(self.accept) <= 1.0: + raise ConfigError("advisor thresholds must satisfy 0 <= review <= accept <= 1") + object.__setattr__( + self, "endpoint_hosts", normalize_names(self.endpoint_hosts, label="advisor.endpoint_hosts") + ) + object.__setattr__( + self, + "allowed_state_fields", + normalize_names(self.allowed_state_fields, label="advisor.allowed_state_fields"), + ) + @dataclass(frozen=True) class OutputConfig: directory: Path @@ -364,6 +393,37 @@ def _instrument_options(raw: dict, section: str) -> dict[str, object]: raise ConfigError(f"{section}.options must be a TOML table") return dict(options) + +def _name_list(raw: dict, key: str, *, path: str) -> tuple[str, ...]: + values = raw.get(key, []) + if isinstance(values, (str, bytes)) or not isinstance(values, (list, tuple)): + raise ConfigError(f"{path}.{key} must be an array of strings") + for value in values: + if not isinstance(value, str) or not value or value.strip() != value: + raise ConfigError(f"{path}.{key} must contain non-empty, trimmed strings") + return tuple(dict.fromkeys(values)) + + +def _advisor_config(raw: object) -> AdvisorConfig: + if not isinstance(raw, dict): + raise ConfigError("advisor must be a TOML table") + enabled = _strict_bool(raw, "enabled", False, path="advisor") + endpoints = _name_list(raw, "endpoint_hosts", path="advisor") + fields = _name_list(raw, "allowed_state_fields", path="advisor") + if enabled and not fields: + raise ConfigError( + "advisor.enabled requires a non-empty advisor.allowed_state_fields allowlist" + ) + if enabled and not endpoints: + raise ConfigError("advisor.enabled requires a non-empty advisor.endpoint_hosts allowlist") + return AdvisorConfig( + enabled=enabled, + endpoint_hosts=endpoints, + allowed_state_fields=fields, + accept=_finite_number(raw.get("accept", 0.60), path="advisor.accept"), + review=_finite_number(raw.get("review", 0.35), path="advisor.review"), + ) + @dataclass(frozen=True) class WaveBenchConfig: connection: ConnectionConfig @@ -380,6 +440,7 @@ class WaveBenchConfig: tui: TuiConfig = TuiConfig() # Append-only: preserve the public positional layout of existing config fields. rf_source: RfSourceConfig | None = None + advisor: AdvisorConfig = AdvisorConfig() def with_connection_timeout_ms(self, timeout_ms: int) -> "WaveBenchConfig": if timeout_ms <= 0: @@ -666,6 +727,7 @@ def load_config(path: str | Path = "wavebench.toml") -> WaveBenchConfig: log_keep_lines_after_trim=int(tui_raw.get("log_keep_lines_after_trim", 1_000)), ), rf_source=rf_source, + advisor=_advisor_config(raw.get("advisor", {})), ) except KeyError as exc: raise ConfigError(f"missing required config key: {exc}") from exc diff --git a/src/wavebench/services/advisor_consent.py b/src/wavebench/services/advisor_consent.py new file mode 100644 index 00000000..969c611d --- /dev/null +++ b/src/wavebench/services/advisor_consent.py @@ -0,0 +1,311 @@ +"""advisor 外部状态外发的同意门(Core 拥有)。 + +契约见 `docs/project/rfcs/WaveBench_advisor插件RFC.md`: + +- 默认关闭,由 `[advisor] enabled` 与调用方的显式 opt-in 共同开启; +- 发送前必须能产出完整预览(payload、字节数、sha256、逐条 untrusted 来源),预览不联网; +- 同意绑定 endpoint 集合与允许字段集,任一变化即失效; +- 非交互场景一律拒绝; +- API key 永不进入 artifact、日志或错误信息。 + +本模块只用标准库,且不含任何网络代码:拒绝路径下调用方没有任何可外发的通道。 +""" + +from __future__ import annotations + +import hashlib +import json +from dataclasses import dataclass +from typing import Any, Collection, Mapping, Sequence + +ADVISOR_CONSENT_SCHEMA = "wavebench.advisor_consent.v1" + +STATUS_GRANTED = "granted" +STATUS_REFUSED = "refused" +STATUS_PREVIEW_ONLY = "preview_only" + +REASON_DISABLED = "advisor_disabled" +REASON_NOT_GRANTED = "external_state_consent_not_granted" +REASON_PREVIEW_ONLY = "preview_only" +REASON_ENDPOINT_NOT_REGISTERED = "endpoint_not_registered" +REASON_FIELD_NOT_REGISTERED = "state_field_not_registered" +REASON_ENDPOINT_CHANGED = "consent_invalidated_by_endpoint_change" +REASON_FIELD_CHANGED = "consent_invalidated_by_allowed_field_change" +REASON_RUN_REQUIRED = "run_scoped_consent_requires_run_target" +REASON_RUN_CHANGED = "consent_invalidated_by_run_change" + +VALID_SCOPES = ("invocation", "run") + + +def normalize_names(values: Sequence[str], *, label: str) -> tuple[str, ...]: + """校验并规格化名字集合(去重、排序),保持产物可比对。""" + if isinstance(values, (str, bytes)): + raise ValueError(f"{label} must be a sequence of names, not a string") + for value in values: + if not isinstance(value, str) or not value or value.strip() != value: + raise ValueError(f"{label} must contain non-empty, trimmed names") + return tuple(sorted(set(values))) + + +@dataclass(frozen=True) +class PreviewPayload: + """即将离开本机的完整请求体;给操作员看,不联网。""" + + payload: Mapping[str, Any] + payload_sha256: str + payload_bytes: int + untrusted_sources: tuple[str, ...] + + def as_dict(self) -> dict[str, Any]: + return { + "payload": dict(self.payload), + "payload_sha256": self.payload_sha256, + "payload_bytes": self.payload_bytes, + "untrusted_sources": list(self.untrusted_sources), + } + + +def build_preview( + *, + model: str, + state: Mapping[str, Any], + questions: Mapping[str, Any], +) -> PreviewPayload: + """构造预览:内容、字节数与 sha256 都由本函数确定,调用方不能跳过。""" + body = {"model": model, "state": dict(state), "questions": dict(questions)} + encoded = json.dumps(body, ensure_ascii=False, sort_keys=True).encode("utf-8") + sources: list[str] = [] + for span in state.get("untrusted", ()) or (): + if isinstance(span, Mapping): + source = str(span.get("source", "")) + if source: + sources.append(source) + return PreviewPayload( + payload=body, + payload_sha256=hashlib.sha256(encoded).hexdigest(), + payload_bytes=len(encoded), + untrusted_sources=tuple(sorted(set(sources))), + ) + + +def unregistered_state_fields(fields: Collection[str], allowed: Collection[str]) -> tuple[str, ...]: + """字段级白名单检查:出现未注册字段就拒绝,避免顺手把新字段送出机器。""" + return tuple(sorted(set(fields) - set(allowed))) + + +@dataclass(frozen=True) +class ConsentDecision: + """操作员对「把这份 state 发到外部服务」的显式决定。 + + 默认按次确认;`valid_for="run"` 时同意只对同一 `run_id` 有效。 + """ + + granted: bool + valid_for: str = "invocation" + run_id: str | None = None + endpoint_hosts: tuple[str, ...] = () + allowed_state_fields: tuple[str, ...] = () + granted_by: str | None = None + granted_at: str | None = None + + def __post_init__(self) -> None: + if self.valid_for not in VALID_SCOPES: + raise ValueError( + f"consent valid_for must be one of {', '.join(VALID_SCOPES)}, got {self.valid_for!r}" + ) + if self.granted and self.valid_for == "run" and not self.run_id: + raise ValueError("run-scoped consent requires run_id") + object.__setattr__( + self, "endpoint_hosts", normalize_names(self.endpoint_hosts, label="endpoint_hosts") + ) + object.__setattr__( + self, + "allowed_state_fields", + normalize_names(self.allowed_state_fields, label="allowed_state_fields"), + ) + + def as_payload(self, preview: PreviewPayload) -> dict[str, Any]: + return { + "granted": self.granted, + "valid_for": self.valid_for, + "run_id": self.run_id, + "granted_by": self.granted_by, + "granted_at": self.granted_at, + "endpoint_hosts": list(self.endpoint_hosts), + "allowed_state_fields": list(self.allowed_state_fields), + "payload_sha256": preview.payload_sha256, + "payload_bytes": preview.payload_bytes, + "accepted_data_leaves_machine": self.granted, + } + + +@dataclass(frozen=True) +class ConsentOutcome: + """同意门的结论;`status != "granted"` 时调用方必须不发起外发。""" + + status: str + reason: str | None + preview: PreviewPayload + decision: ConsentDecision | None = None + endpoint_hosts: tuple[str, ...] = () + allowed_state_fields: tuple[str, ...] = () + + @property + def granted(self) -> bool: + return self.status == STATUS_GRANTED + + def as_payload(self) -> dict[str, Any]: + record: dict[str, Any] = { + "schema": ADVISOR_CONSENT_SCHEMA, + "status": self.status, + "reason": self.reason, + "endpoint_hosts": list(self.endpoint_hosts), + "allowed_state_fields": list(self.allowed_state_fields), + "payload_sha256": self.preview.payload_sha256, + "payload_bytes": self.preview.payload_bytes, + "accepted_data_leaves_machine": self.granted, + } + if self.decision is not None: + record.update( + { + "granted": self.decision.granted, + "valid_for": self.decision.valid_for, + "run_id": self.decision.run_id, + "granted_by": self.decision.granted_by, + "granted_at": self.decision.granted_at, + } + ) + else: + record.update({"granted": False, "valid_for": None, "run_id": None, + "granted_by": None, "granted_at": None}) + return record + + +def consent_covers( + consent: ConsentDecision, + *, + endpoint_hosts: Sequence[str], + allowed_state_fields: Sequence[str], + run_id: str | None, +) -> tuple[bool, str | None]: + """判断已有同意是否覆盖本次调用;返回 (是否覆盖, 失效原因)。""" + if not consent.granted: + return False, REASON_NOT_GRANTED + if set(consent.endpoint_hosts) != set(endpoint_hosts): + return False, REASON_ENDPOINT_CHANGED + if set(consent.allowed_state_fields) != set(allowed_state_fields): + return False, REASON_FIELD_CHANGED + if consent.valid_for == "invocation": + return True, None + if run_id is None: + return False, REASON_RUN_REQUIRED + if consent.run_id != run_id: + return False, REASON_RUN_CHANGED + return True, None + + +def resolve_consent( + *, + preview: PreviewPayload, + endpoint_hosts: Sequence[str], + state_fields: Collection[str], + enabled: bool, + registered_endpoint_hosts: Sequence[str] = (), + registered_state_fields: Collection[str] = (), + run_id: str | None = None, + preview_only: bool = False, + accepted: bool = False, + valid_for: str = "invocation", + granted_by: str | None = None, + granted_at: str | None = None, +) -> ConsentOutcome: + """Core 侧的同意门。 + + 拒绝顺序(对应 RFC 验收门里"外部调用次数必须为 0"的五种情况):未开启 → 未注册字段 → + 未注册 endpoint → 仅预览 → 未确认。"仅预览"与"拒绝"都不返回可用同意,调用方不得外发。 + """ + hosts = normalize_names(endpoint_hosts, label="endpoint_hosts") + fields = normalize_names(tuple(state_fields), label="state_fields") + + def refused(reason: str) -> ConsentOutcome: + return ConsentOutcome( + status=STATUS_REFUSED, + reason=reason, + preview=preview, + endpoint_hosts=hosts, + allowed_state_fields=fields, + ) + + if not enabled: + return refused(REASON_DISABLED) + + unregistered = unregistered_state_fields(fields, registered_state_fields) + if registered_state_fields and unregistered: + return refused(f"{REASON_FIELD_NOT_REGISTERED}: {', '.join(unregistered)}") + + if registered_endpoint_hosts: + allowed_hosts = set(normalize_names(registered_endpoint_hosts, label="registered_endpoint_hosts")) + unknown_hosts = tuple(sorted(set(hosts) - allowed_hosts)) + if unknown_hosts: + return refused(f"{REASON_ENDPOINT_NOT_REGISTERED}: {', '.join(unknown_hosts)}") + + if preview_only: + return ConsentOutcome( + status=STATUS_PREVIEW_ONLY, + reason=REASON_PREVIEW_ONLY, + preview=preview, + endpoint_hosts=hosts, + allowed_state_fields=fields, + ) + + if not accepted: + return refused(REASON_NOT_GRANTED) + + if valid_for not in VALID_SCOPES: + return refused(REASON_NOT_GRANTED) + if valid_for == "run" and not run_id: + return refused(REASON_RUN_REQUIRED) + + decision = ConsentDecision( + granted=True, + valid_for=valid_for, + run_id=run_id if valid_for == "run" else None, + endpoint_hosts=hosts, + allowed_state_fields=fields, + granted_by=granted_by, + granted_at=granted_at, + ) + return ConsentOutcome( + status=STATUS_GRANTED, + reason=None, + preview=preview, + decision=decision, + endpoint_hosts=hosts, + allowed_state_fields=fields, + ) + + +__all__ = [ + "ADVISOR_CONSENT_SCHEMA", + "REASON_DISABLED", + "REASON_ENDPOINT_CHANGED", + "REASON_ENDPOINT_NOT_REGISTERED", + "REASON_FIELD_CHANGED", + "REASON_FIELD_NOT_REGISTERED", + "REASON_NOT_GRANTED", + "REASON_PREVIEW_ONLY", + "REASON_RUN_CHANGED", + "REASON_RUN_REQUIRED", + "STATUS_GRANTED", + "STATUS_PREVIEW_ONLY", + "STATUS_REFUSED", + "VALID_SCOPES", + "ConsentDecision", + "ConsentOutcome", + "PreviewPayload", + "build_preview", + "consent_covers", + "normalize_names", + "resolve_consent", + "unregistered_state_fields", +] diff --git a/src/wavebench/services/decision_artifacts.py b/src/wavebench/services/decision_artifacts.py new file mode 100644 index 00000000..4ece24e4 --- /dev/null +++ b/src/wavebench/services/decision_artifacts.py @@ -0,0 +1,179 @@ +"""advisor decision artifact(Core 拥有):schema `wavebench.decision.v1`。 + +契约见 `docs/project/rfcs/WaveBench_advisor插件RFC.md`: + +- 写入对应 run 目录的 `decisions/`,附加式,绝不修改 `run.json` / `summary.csv` / `steps/*`; +- 文件名含 UTC 时间戳与 advisor id,独占创建; +- 无 run 目录时不落盘; +- "概率 → 动作"的阈值策略由 Core 配置拥有,插件不得自带阈值; +- 产物只做建议,不参与 `run.json.status`、质量门、`auto_recover` 或 capability 判定。 +""" + +from __future__ import annotations + +import json +import re +from dataclasses import dataclass, field +from datetime import UTC, datetime +from pathlib import Path +from typing import Any, Mapping + +ADVISOR_SCHEMA = "wavebench.decision.v1" +DECISIONS_DIRNAME = "decisions" + +ARTIFACT_STATUSES = ("ok", "refused", "preview_only", "unavailable", "invalid") +# 文件名安全:只允许字母、数字、下划线、点与连字符(冒号在 Windows 上是非法字符)。 +_SAFE_TOKEN = re.compile(r"^[A-Za-z0-9_.-]{1,96}$") + + +class DecisionArtifactError(ValueError): + """artifact 契约被违反(不是网络或服务错误)。""" + + +@dataclass(frozen=True) +class ThresholdPolicy: + """概率 → 动作的阈值;低于 accept 给人工复核,低于 review 不给建议。""" + + accept: float = 0.60 + review: float = 0.35 + + def __post_init__(self) -> None: + for label, value in (("accept", self.accept), ("review", self.review)): + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise DecisionArtifactError(f"threshold {label} must be a number") + if not 0.0 <= float(value) <= 1.0: + raise DecisionArtifactError(f"threshold {label} must be within 0..1") + if self.review > self.accept: + raise DecisionArtifactError("review threshold must not exceed accept threshold") + + def as_record(self) -> dict[str, float]: + return {"accept": float(self.accept), "review": float(self.review)} + + +@dataclass(frozen=True) +class DecisionTarget: + """artifact 的落点:对应 run 目录,以及可选的 run.json 指纹。""" + + run_dir: Path + run_json_sha256: str | None = None + + @property + def run_id(self) -> str: + return Path(self.run_dir).name + + def as_payload(self) -> dict[str, Any]: + return { + "run_id": self.run_id, + "run_dir": str(self.run_dir).replace("\\", "/"), + "run_json_sha256": self.run_json_sha256, + } + + +@dataclass(frozen=True) +class DecisionArtifact: + """`wavebench.decision.v1` 的正文;任何候选结论都只作为建议留痕。""" + + status: str + requested_model: str + reason: str | None = None + reported_model: str | None = None + target: Mapping[str, Any] | None = None + consent: Mapping[str, Any] | None = None + state: Mapping[str, Any] = field(default_factory=dict) + questions: Mapping[str, Any] = field(default_factory=dict) + answers: Mapping[str, Any] = field(default_factory=dict) + thresholds: Mapping[str, Any] = field(default_factory=dict) + recommendations: tuple[Mapping[str, Any], ...] = () + duration_ms: int = 0 + schema: str = ADVISOR_SCHEMA + advisory_only: bool = True + + def __post_init__(self) -> None: + if self.status not in ARTIFACT_STATUSES: + raise DecisionArtifactError( + f"artifact status must be one of {', '.join(ARTIFACT_STATUSES)}, got {self.status!r}" + ) + if not self.advisory_only: + raise DecisionArtifactError("decision artifacts must stay advisory-only") + if isinstance(self.duration_ms, bool) or not isinstance(self.duration_ms, int): + raise DecisionArtifactError("duration_ms must be an integer") + if self.duration_ms < 0: + raise DecisionArtifactError("duration_ms must be >= 0") + + def as_dict(self) -> dict[str, Any]: + return { + "schema": self.schema, + "advisory_only": self.advisory_only, + "status": self.status, + "reason": self.reason, + "target": dict(self.target) if self.target is not None else None, + "advisory": { + "requested_model": self.requested_model, + "reported_model": self.reported_model, + "duration_ms": self.duration_ms, + }, + "consent": dict(self.consent) if self.consent is not None else None, + "state": dict(self.state), + "questions": dict(self.questions), + "answers": dict(self.answers), + "thresholds": dict(self.thresholds), + "recommendations": [dict(item) for item in self.recommendations], + } + + def to_json(self) -> str: + return json.dumps(self.as_dict(), indent=2, ensure_ascii=False, sort_keys=True) + + +def utc_stamp(moment: datetime | None = None) -> str: + """文件名用的 UTC 时间戳(无冒号,Windows 路径安全)。""" + return (moment or datetime.now(UTC)).astimezone(UTC).strftime("%Y%m%dT%H%M%SZ") + + +def decisions_dir(run_dir: str | Path) -> Path: + return Path(run_dir) / DECISIONS_DIRNAME + + +def artifact_path(run_dir: str | Path, *, advisor_id: str, timestamp: str) -> Path: + for label, value in (("advisor_id", advisor_id), ("timestamp", timestamp)): + if not isinstance(value, str) or not _SAFE_TOKEN.match(value): + raise DecisionArtifactError( + f"{label} must match {_SAFE_TOKEN.pattern} to stay a safe file name" + ) + return decisions_dir(run_dir) / f"{timestamp}-{advisor_id}.json" + + +def write_artifact( + artifact: DecisionArtifact, + run_dir: str | Path, + *, + advisor_id: str, + timestamp: str | None = None, +) -> Path | None: + """独占创建 artifact;无 run 目录时不落盘(返回 None)。 + + 产物永远是新文件:不覆盖任何既有文件,也不触碰 Core 拥有的 `run.json` / `summary.csv` / + `steps/*`,因此不需要走原子替换路径;文件名带 UTC 时间戳,重名即视为编程错误。 + """ + run = Path(run_dir) + if not run.is_dir(): + return None + path = artifact_path(run, advisor_id=advisor_id, timestamp=timestamp or utc_stamp()) + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("x", encoding="utf-8") as handle: + handle.write(artifact.to_json()) + return path + + +__all__ = [ + "ADVISOR_SCHEMA", + "ARTIFACT_STATUSES", + "DECISIONS_DIRNAME", + "DecisionArtifact", + "DecisionArtifactError", + "DecisionTarget", + "ThresholdPolicy", + "artifact_path", + "decisions_dir", + "utc_stamp", + "write_artifact", +] diff --git a/tests/test_advisor_config.py b/tests/test_advisor_config.py new file mode 100644 index 00000000..6b68206e --- /dev/null +++ b/tests/test_advisor_config.py @@ -0,0 +1,75 @@ +"""`[advisor]` 配置段的解析与校验测试。""" + +from __future__ import annotations + +import tempfile +import unittest +from pathlib import Path + +from wavebench.config import load_config +from wavebench.errors import ConfigError + +BASE = """\ +[connection] +resource = "TCPIP::192.0.2.40::INSTR" +[scope] +""" + + +class AdvisorConfigTests(unittest.TestCase): + def _load(self, advisor_section: str): + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / "wavebench.toml" + path.write_text(BASE + advisor_section, encoding="utf-8") + return load_config(path) + + def test_defaults_to_disabled_without_allowlists(self): + config = self._load("") + + self.assertFalse(config.advisor.enabled) + self.assertEqual(config.advisor.endpoint_hosts, ()) + self.assertEqual(config.advisor.allowed_state_fields, ()) + self.assertEqual(config.advisor.accept, 0.60) + self.assertEqual(config.advisor.review, 0.35) + + def test_enabled_config_normalizes_allowlists(self): + config = self._load( + """ +[advisor] +enabled = true +endpoint_hosts = ["api.typesafe.ai", "api.typesafe.ai"] +allowed_state_fields = ["run_status", "cycles_bucket"] +accept = 0.8 +review = 0.4 +""" + ) + + self.assertTrue(config.advisor.enabled) + self.assertEqual(config.advisor.endpoint_hosts, ("api.typesafe.ai",)) + self.assertEqual(config.advisor.allowed_state_fields, ("cycles_bucket", "run_status")) + self.assertEqual(config.advisor.accept, 0.8) + self.assertEqual(config.advisor.review, 0.4) + + def test_enabled_requires_allowlists(self): + with self.assertRaisesRegex(ConfigError, "advisor.allowed_state_fields"): + self._load("\n[advisor]\nenabled = true\nendpoint_hosts = [\"api.typesafe.ai\"]\n") + with self.assertRaisesRegex(ConfigError, "advisor.endpoint_hosts"): + self._load("\n[advisor]\nenabled = true\nallowed_state_fields = [\"run_status\"]\n") + + def test_threshold_order_is_enforced(self): + with self.assertRaisesRegex(ConfigError, "0 <= review <= accept <= 1"): + self._load( + "\n[advisor]\naccept = 0.2\nreview = 0.9\n" + ) + + def test_endpoint_hosts_must_be_an_array_of_strings(self): + with self.assertRaisesRegex(ConfigError, "endpoint_hosts must be an array of strings"): + self._load("\n[advisor]\nendpoint_hosts = \"api.typesafe.ai\"\n") + + def test_unknown_boolean_value_is_rejected(self): + with self.assertRaisesRegex(ConfigError, "advisor.enabled must be a boolean"): + self._load("\n[advisor]\nenabled = \"yes\"\n") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_advisor_consent.py b/tests/test_advisor_consent.py new file mode 100644 index 00000000..28fdf93f --- /dev/null +++ b/tests/test_advisor_consent.py @@ -0,0 +1,214 @@ +"""advisor 同意门的离线契约测试(无网络、无 API key、无第三方 SDK)。""" + +from __future__ import annotations + +import ast +import json +import unittest +from pathlib import Path + +from wavebench.services import advisor_consent as consent + +MODULE_PATH = Path(consent.__file__) +ALLOWED_IMPORTS = {"__future__", "dataclasses", "hashlib", "json", "typing"} + +STATE = {"fields": {"run_status": "failed", "cycles_bucket": "few_cycles"}} +QUESTIONS = {"route": {"kind": "choice", "options": ["run check", "run verify"]}} + + +class AdvisorConsentTests(unittest.TestCase): + def _preview(self, *, state=None, questions=None): + return consent.build_preview( + model="jev-1.13.0", + state=state if state is not None else STATE, + questions=questions if questions is not None else QUESTIONS, + ) + + def _resolve(self, *, preview=None, **overrides): + kwargs = { + "preview": preview if preview is not None else self._preview(), + "endpoint_hosts": ("api.typesafe.ai",), + "state_fields": ("run_status", "cycles_bucket"), + "enabled": True, + "registered_endpoint_hosts": ("api.typesafe.ai",), + "registered_state_fields": ("run_status", "cycles_bucket"), + "accepted": True, + } + kwargs.update(overrides) + return consent.resolve_consent(**kwargs) + + def test_preview_is_deterministic_and_counts_untrusted_sources(self): + state = { + "fields": {"run_status": "failed"}, + "untrusted": [ + {"source": "instrument.response", "text": "IDN: ..."}, + {"source": "operator.utterance", "text": "信号不对"}, + {"source": "instrument.response", "text": "dup"}, + ], + } + first = self._preview(state=state) + second = self._preview(state=state) + + self.assertEqual(first.payload_sha256, second.payload_sha256) + self.assertEqual(first.payload_bytes, len(json.dumps(first.payload, ensure_ascii=False, + sort_keys=True).encode("utf-8"))) + self.assertEqual(first.untrusted_sources, ("instrument.response", "operator.utterance")) + + def test_preview_changes_when_payload_changes(self): + self.assertNotEqual( + self._preview().payload_sha256, + self._preview(state={"fields": {"run_status": "ok"}}).payload_sha256, + ) + + def test_disabled_gate_refuses(self): + outcome = self._resolve(enabled=False) + + self.assertEqual(outcome.status, consent.STATUS_REFUSED) + self.assertEqual(outcome.reason, consent.REASON_DISABLED) + self.assertIsNone(outcome.decision) + self.assertFalse(outcome.granted) + + def test_unregistered_state_field_refuses_with_field_name(self): + outcome = self._resolve(state_fields=("run_status", "cycles_bucket", "raw_voltage_v")) + + self.assertEqual(outcome.reason, "state_field_not_registered: raw_voltage_v") + self.assertIsNone(outcome.decision) + + def test_unregistered_endpoint_refuses_with_host_name(self): + outcome = self._resolve(endpoint_hosts=("api.typesafe.ai", "evil.example")) + + self.assertEqual(outcome.reason, "endpoint_not_registered: evil.example") + self.assertIsNone(outcome.decision) + + def test_preview_only_never_grants_consent(self): + outcome = self._resolve(preview_only=True, accepted=True) + + self.assertEqual(outcome.status, consent.STATUS_PREVIEW_ONLY) + self.assertEqual(outcome.reason, consent.REASON_PREVIEW_ONLY) + self.assertIsNone(outcome.decision) + self.assertFalse(outcome.granted) + + def test_missing_confirmation_refuses(self): + outcome = self._resolve(accepted=False) + + self.assertEqual(outcome.reason, consent.REASON_NOT_GRANTED) + self.assertIsNone(outcome.decision) + + def test_granted_consent_records_scope_and_allowlists(self): + outcome = self._resolve(granted_by="operator", granted_at="2026-09-29T10:00:00Z") + + self.assertTrue(outcome.granted) + self.assertIsNotNone(outcome.decision) + assert outcome.decision is not None + self.assertEqual(outcome.decision.valid_for, "invocation") + self.assertIsNone(outcome.decision.run_id) + self.assertEqual(outcome.decision.endpoint_hosts, ("api.typesafe.ai",)) + self.assertEqual( + outcome.decision.allowed_state_fields, ("cycles_bucket", "run_status") + ) + self.assertEqual(outcome.as_payload()["granted_by"], "operator") + + def test_run_scoped_consent_requires_run_target(self): + outcome = self._resolve(valid_for="run", run_id=None) + + self.assertEqual(outcome.reason, consent.REASON_RUN_REQUIRED) + self.assertIsNone(outcome.decision) + + def test_run_scoped_consent_is_bound_to_run_id(self): + granted = self._resolve(valid_for="run", run_id="20260929_0740_loop_gain") + assert granted.decision is not None + self.assertEqual(granted.decision.run_id, "20260929_0740_loop_gain") + + covered, reason = consent.consent_covers( + granted.decision, + endpoint_hosts=("api.typesafe.ai",), + allowed_state_fields=("cycles_bucket", "run_status"), + run_id="20260929_0800_other_run", + ) + self.assertFalse(covered) + self.assertEqual(reason, consent.REASON_RUN_CHANGED) + + def test_consent_invalidated_by_endpoint_or_field_change(self): + granted = self._resolve() + assert granted.decision is not None + + covered, reason = consent.consent_covers( + granted.decision, + endpoint_hosts=("api.typesafe.ai", "cdn.typesafe.ai"), + allowed_state_fields=("cycles_bucket", "run_status"), + run_id=None, + ) + self.assertFalse(covered) + self.assertEqual(reason, consent.REASON_ENDPOINT_CHANGED) + + covered, reason = consent.consent_covers( + granted.decision, + endpoint_hosts=("api.typesafe.ai",), + allowed_state_fields=("run_status",), + run_id=None, + ) + self.assertFalse(covered) + self.assertEqual(reason, consent.REASON_FIELD_CHANGED) + + def test_invocation_consent_covers_repeated_calls(self): + granted = self._resolve() + assert granted.decision is not None + + covered, reason = consent.consent_covers( + granted.decision, + endpoint_hosts=("api.typesafe.ai",), + allowed_state_fields=("cycles_bucket", "run_status"), + run_id=None, + ) + self.assertTrue(covered) + self.assertIsNone(reason) + + def test_ungranted_decision_is_not_reusable(self): + decision = consent.ConsentDecision(granted=False, endpoint_hosts=("api.typesafe.ai",)) + + covered, reason = consent.consent_covers( + decision, + endpoint_hosts=("api.typesafe.ai",), + allowed_state_fields=(), + run_id=None, + ) + self.assertFalse(covered) + self.assertEqual(reason, consent.REASON_NOT_GRANTED) + + def test_consent_record_contains_required_keys_without_secrets(self): + record = self._resolve().as_payload() + + for key in ( + "granted", + "valid_for", + "run_id", + "granted_by", + "granted_at", + "endpoint_hosts", + "allowed_state_fields", + "payload_sha256", + "payload_bytes", + "accepted_data_leaves_machine", + ): + self.assertIn(key, record) + self.assertNotIn("api_key", json.dumps(record).lower()) + + def test_module_has_no_network_or_third_party_imports(self): + tree = ast.parse(MODULE_PATH.read_text(encoding="utf-8")) + imported: set[str] = set() + for node in ast.walk(tree): + if isinstance(node, ast.Import): + imported.update(alias.name.split(".")[0] for alias in node.names) + elif isinstance(node, ast.ImportFrom) and node.module: + imported.add(node.module.split(".")[0]) + + self.assertTrue( + imported <= ALLOWED_IMPORTS, + f"unexpected imports in advisor_consent: {sorted(imported - ALLOWED_IMPORTS)}", + ) + for banned in ("socket", "urllib", "http", "requests", "httpx", "typesafe_sdk"): + self.assertNotIn(banned, imported) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_decision_artifacts.py b/tests/test_decision_artifacts.py new file mode 100644 index 00000000..00b7e052 --- /dev/null +++ b/tests/test_decision_artifacts.py @@ -0,0 +1,125 @@ +"""decision artifact(`wavebench.decision.v1`)的离线契约测试。""" + +from __future__ import annotations + +import json +import re +import tempfile +import unittest +from pathlib import Path + +from wavebench.services import decision_artifacts as artifacts + + +def _artifact(**overrides) -> artifacts.DecisionArtifact: + kwargs = { + "status": "refused", + "requested_model": "jev-1.13.0", + "reason": "external_state_consent_not_granted", + "consent": {"schema": "wavebench.advisor_consent.v1", "granted": False}, + "state": {"fields": {"run_status": "failed"}}, + "thresholds": artifacts.ThresholdPolicy().as_record(), + } + kwargs.update(overrides) + return artifacts.DecisionArtifact(**kwargs) + + +class DecisionArtifactTests(unittest.TestCase): + def test_threshold_policy_bounds(self): + self.assertEqual(artifacts.ThresholdPolicy().as_record(), {"accept": 0.60, "review": 0.35}) + with self.assertRaises(artifacts.DecisionArtifactError): + artifacts.ThresholdPolicy(accept=0.3, review=0.5) + with self.assertRaises(artifacts.DecisionArtifactError): + artifacts.ThresholdPolicy(accept=1.5) + + def test_artifact_rejects_unknown_status_and_silent_downgrade(self): + with self.assertRaises(artifacts.DecisionArtifactError): + _artifact(status="done") + with self.assertRaises(artifacts.DecisionArtifactError): + _artifact(advisory_only=False) + with self.assertRaises(artifacts.DecisionArtifactError): + _artifact(duration_ms=-1) + + def test_artifact_payload_keeps_schema_and_recommendations(self): + payload = _artifact( + status="ok", + reason=None, + reported_model="jev-1.13.0", + duration_ms=214, + recommendations=({"action": "run check", "probability": 0.81},), + ).as_dict() + + self.assertEqual(payload["schema"], "wavebench.decision.v1") + self.assertTrue(payload["advisory_only"]) + self.assertEqual(payload["advisory"]["requested_model"], "jev-1.13.0") + self.assertEqual(payload["advisory"]["reported_model"], "jev-1.13.0") + self.assertEqual(payload["advisory"]["duration_ms"], 214) + self.assertEqual(payload["recommendations"][0]["action"], "run check") + + def test_artifact_path_layout_and_token_safety(self): + run_dir = Path("data") / "runs" / "20260929_0740_loop_gain" + path = artifacts.artifact_path( + run_dir, advisor_id="typesafe.jev", timestamp="20260929T104827Z" + ) + + self.assertEqual( + path.as_posix(), + "data/runs/20260929_0740_loop_gain/decisions/20260929T104827Z-typesafe.jev.json", + ) + with self.assertRaises(artifacts.DecisionArtifactError): + artifacts.artifact_path(run_dir, advisor_id="../evil", timestamp="20260929T104827Z") + with self.assertRaises(artifacts.DecisionArtifactError): + artifacts.artifact_path(run_dir, advisor_id="typesafe.jev", timestamp="20260929:10:48") + + def test_utc_stamp_is_file_name_safe(self): + self.assertRegex(artifacts.utc_stamp(), re.compile(r"^\d{8}T\d{6}Z$")) + + def test_write_artifact_is_exclusive_and_leaves_core_files_untouched(self): + with tempfile.TemporaryDirectory() as tmp: + run_dir = Path(tmp) / "20260929_0740_loop_gain" + steps_dir = run_dir / "steps" + steps_dir.mkdir(parents=True) + run_json = run_dir / "run.json" + summary = run_dir / "summary.csv" + step_file = steps_dir / "step1.json" + run_json.write_text('{"status": "failed"}', encoding="utf-8") + summary.write_text("step,status\n1,failed\n", encoding="utf-8") + step_file.write_text('{"step": 1}', encoding="utf-8") + before = {path: path.read_bytes() for path in (run_json, summary, step_file)} + + written = artifacts.write_artifact( + _artifact(), run_dir, advisor_id="typesafe.jev", timestamp="20260929T104827Z" + ) + + self.assertEqual(written, run_dir / "decisions" / "20260929T104827Z-typesafe.jev.json") + payload = json.loads(written.read_text(encoding="utf-8")) + self.assertEqual(payload["schema"], "wavebench.decision.v1") + self.assertEqual(payload["status"], "refused") + for path, content in before.items(): + self.assertEqual(path.read_bytes(), content, f"{path.name} must not change") + self.assertEqual(sorted(p.name for p in run_dir.iterdir()), ["decisions", "run.json", "steps", "summary.csv"]) + + with self.assertRaises(FileExistsError): + artifacts.write_artifact( + _artifact(), run_dir, advisor_id="typesafe.jev", timestamp="20260929T104827Z" + ) + + def test_write_artifact_without_run_dir_writes_nothing(self): + with tempfile.TemporaryDirectory() as tmp: + missing = Path(tmp) / "not_a_run" + + written = artifacts.write_artifact( + _artifact(), missing, advisor_id="typesafe.jev", timestamp="20260929T104827Z" + ) + + self.assertIsNone(written) + self.assertFalse(missing.exists()) + + def test_decisions_dir_is_relative_to_run(self): + self.assertEqual( + artifacts.decisions_dir(Path("data/runs/x")).as_posix(), "data/runs/x/decisions" + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/wavebench.example.toml b/wavebench.example.toml index d83fadb8..986f59d0 100644 --- a/wavebench.example.toml +++ b/wavebench.example.toml @@ -262,3 +262,18 @@ settle_ms_before_read = 0 # Extra wait after changing DMM function before the first read. # 切换万用表功能后、第一次读数前的额外等待时间。 settle_ms_after_function_change = 500 + +[advisor] +# advisor(外部判断模型,例如 TypeSafe Jev)总开关。默认 false:关闭时不会构造任何外发请求。 +# 打开时必须同时给出 endpoint_hosts 与 allowed_state_fields 白名单,否则配置校验直接失败。 +enabled = false + +# 允许外发的服务主机白名单(严格匹配)。留空表示不允许任何 endpoint。 +endpoint_hosts = [] +# 允许外发的状态字段白名单(严格匹配,字段名与插件声明一致)。留空表示不允许任何字段。 +allowed_state_fields = [] + +# "概率 -> 动作"的阈值由 Core 拥有,插件不得自带;需满足 0 <= review <= accept <= 1。 +accept = 0.60 +review = 0.35 +