diff --git a/docs/how-to/serve-mcp.md b/docs/how-to/serve-mcp.md index 98b81a7f..02f9437f 100644 --- a/docs/how-to/serve-mcp.md +++ b/docs/how-to/serve-mcp.md @@ -1,6 +1,6 @@ # 启动只读 MCP 服务 -WaveBench HTTP MCP 服务提供本机或受控网络中的离线信息和 run plan 检查。它不提供 raw SCPI、输出控制或 run 执行。 +WaveBench HTTP MCP 服务提供本机或受控网络中的离线信息和只读仪器观察。它不提供 raw SCPI、输出控制或 run 执行。 ## 启动服务 @@ -20,7 +20,23 @@ python -m wavebench mcp serve \ | `GET /tools` | Bearer token | 列出只读工具。 | | `POST /call`、`POST /mcp` | Bearer token | 调用 MCP 工具或 JSON-RPC 方法。 | -请求体上限为 1 MiB。当前工具仅包含 `run.schema`、`run.check` 和 `capture.inspect`;它们不会连接仪器。`run.check` 只接受项目内的 `plans/*.toml`,`capture.inspect` 只读取项目内的离线采集包。 +请求体上限为 1 MiB。所有工具都是只读的,不会改变仪器状态,也不会读取波形。工具的权威列表和元数据以 `GET /tools` 返回为准: + +| 工具 | 行为 | +| --- | --- | +| `run.schema` | 返回 run plan schema。 | +| `run.check` | 只接受项目内的 `plans/*.toml`,离线解析并校验 run plan。 | +| `capture.inspect` | 只读取项目内的离线采集包摘要。 | +| `doctor.config` | 对配置中的仪器执行只读 doctor 检查,返回结构化记录。 | +| `scope.observe` | 读取示波器身份、每通道状态快照和输入耦合安全;不读取波形。 | +| `scope.advise` | 基于只读状态快照和调用方给出的 `expected_frequencies_hz` 建议显示/采集参数;不应用建议。 | + +`doctor.config`、`scope.observe` 和 `scope.advise` 是实验性工具:开发线已实现并有离线测试,但尚未随正式版本发布, +支持范围不作承诺。 + +`scope.observe` 和 `scope.advise` 不读取示波器波形。波形摘要、期望值检查和跨通道关系分析需要 +先显式读取波形(属于会改变仪器状态的写路径),请在操作者明确执行 +[`scope observe --fetch-waveform`](../reference/cli.md) 时进行,不通过 MCP 暴露。 ## Verification diff --git a/docs/reference/cli.md b/docs/reference/cli.md index de43f36f..c5362125 100644 --- a/docs/reference/cli.md +++ b/docs/reference/cli.md @@ -18,11 +18,41 @@ python -m wavebench --json ... | --- | --- | --- | | 离线、只读本地文件 | `run schema`、`run check`、`run intent`、`run compare`、`run resume`、`capture inspect`、`capability explain`、`lock status` | 不连接仪器。 | | 离线且可能写本地文件 | `run template --output`、`run report`、`run report-index` | 不连接仪器,但会创建模板、报告或索引文件。 | -| 连接读取或预检 | `doctor`、状态/身份查询、`run verify` | 会访问配置的仪器,不应改变实验设置。 | -| 可能改变状态或触发采集 | 输出和 setter、`scope auto`/`scope capture`、非 fake TUI、`run plan` | 可能写入仪器、触发采集或切换输出。 | +| 连接读取或预检 | `doctor`、状态/身份查询、`scope observe`(不带 `--fetch-waveform`)、`run verify` | 会访问配置的仪器,不应改变实验设置。 | +| 可能改变状态或触发采集 | 输出和 setter、`scope auto`/`scope capture`、`scope observe --fetch-waveform`、非 fake TUI、`run plan` | 可能写入仪器、触发采集或切换输出。 | `scope fetch` 读取已有波形,但仍是仪器 I/O;不要把它当作离线命令。每次硬件操作前确认接线、输入阻抗、输出状态和安全限制。WaveBench 不会自动执行 `*RST`,也不会因设置电压、幅度或频率而自动开启输出。 +`scope observe` 是实验性命令:开发线已实现并有离线测试,但尚未随正式版本发布,兼容性和支持范围不作承诺。 +它默认只读:只查询身份、通道状态快照和输入耦合安全,不读取波形。 +`scope observe --fetch-waveform` 是显式写路径,可能停止正在运行的采集、修改波形传输 +source/mode/format/points 并打开通道显示;它逐通道读取波形,因此多通道结果不保证来自同一次 +acquisition,此时跨通道的相位、相关性、延迟和交点不会被计算(`correlation`、`intersections` +返回 `skipped`,相位为 `null`)。需要驱动可证明的同一次采集时,使用 +`scope capture --synchronized`(通道和输出格式要求见其 `--help`)。 + +期望值检查通过 `--expect ` 提供,需要 `--fetch-waveform`: + +```toml +[channels.1] +frequency_hz = 1000 +frequency_tolerance_ratio = 0.05 +vpp_v = 3.3 +duty_percent = 50 +``` + +字段名、类型、有限性和取值范围由实现严格校验(`src/wavebench/data/expectations.py` 的 `validate_expectation()`); +拼错字段名会直接报错,不会被静默忽略。上面的 TOML 只示范格式,字段全集以该实现为准。任何输入错误都在 +打开仪器会话之前被拒绝,因此不会产生仪器写入。 + +`--target-cycles` 和 `--target-vertical-divisions` 必须为有限正数,并在加载配置或创建仪器服务 +之前校验。生成的 focus 建议使用 `--vertical-scale CHANNEL=V_PER_DIV`;隐藏其他通道的参数为 +`--hide-others`。建议不会自动执行。 + +期望值汇总的 `channels` 保留所有待验收通道。波形读取、安全检查或期望值计算失败的通道标记为 +`unavailable`;没有可用检查结果时汇总为 `unavailable`,部分通道已有 `pass`/`warn` 结果时为 +`partial`。已确认的 `fail` 仍优先返回 `fail`。未提供期望值或期望值没有可执行指标时保持 `skipped`。 + ## JSON 输出与退出码 将 `--json` 放在命令行任意位置可请求机器可读输出。成功结果使用 `wavebench.cli.result.v1`,包含 `status`、`exit_code` 和 `result`;错误使用 `wavebench.error.v1`。普通成功输出写入标准输出,普通错误写入标准错误。 diff --git a/src/wavebench/cli.py b/src/wavebench/cli.py index dfd2e5d4..ceb10428 100644 --- a/src/wavebench/cli.py +++ b/src/wavebench/cli.py @@ -10,6 +10,8 @@ from pathlib import Path import sys from tempfile import TemporaryFile +import tomllib +from typing import Any import numpy as np @@ -113,6 +115,8 @@ from .plugins.registry import build_plugin_registry, has_doctor_errors, plugin_doctor_records from .plugins.scpi import has_scpi_doctor_errors, load_scpi_plugin, probe_scpi_plugin, scpi_plugin_doctor_records from .services.scope_service import ScopeService +from .services.agent_observe import scope_observe_payload, scope_waveform_report_payload +from .services.agent_advise import scope_advise_from_observation, validate_scope_advice_targets from .services.source_service import SourceService from .services.rf_source_service import RfSourceService from .services.power_service import PowerService @@ -611,6 +615,137 @@ def _scope_error_check(args: argparse.Namespace) -> ErrorCheckSpec | None: raise ConfigError(str(exc)) from exc +def _load_scope_expectations(path: str | None) -> dict[int, dict[str, Any]] | None: + """读取 --expect 指定的 TOML 期望值文件;校验在读取阶段完成,早于任何仪器 I/O。""" + if path is None: + return None + expectation_path = Path(path) + try: + raw = tomllib.loads(expectation_path.read_bytes().decode("utf-8-sig")) + except OSError as exc: + raise ConfigError(f"failed to read scope expectation file: {expectation_path}") from exc + except tomllib.TOMLDecodeError as exc: + raise ConfigError(f"invalid scope expectation TOML: {expectation_path}") from exc + unknown = sorted(set(raw) - {"channels"}) + if unknown: + raise ConfigError(f"unknown scope expectation section(s): {', '.join(unknown)}") + channels = raw.get("channels") + if not isinstance(channels, dict) or not channels: + raise ConfigError("scope expectation file must define a non-empty [channels] table") + expectations: dict[int, dict[str, Any]] = {} + for key, value in channels.items(): + try: + channel = int(key) + except (TypeError, ValueError) as exc: + raise ConfigError("scope expectation channel keys must be numbers") from exc + if channel < 1: + raise ConfigError("scope expectation channels must be >= 1") + if not isinstance(value, dict): + raise ConfigError("scope expectation channel entries must be tables") + expectations[channel] = dict(value) + return expectations + + +def _expectation_frequencies( + expectations: dict[int, dict[str, Any]] | None, +) -> dict[int, float] | None: + if not expectations: + return None + values: dict[int, float] = {} + for channel, expectation in expectations.items(): + value = expectation.get("frequency_hz") + if isinstance(value, (int, float)) and not isinstance(value, bool) and value > 0: + values[channel] = float(value) + return values or None + + +def _run_scope_observe(args: argparse.Namespace) -> dict[str, Any]: + target_cycles, target_vertical_divisions = validate_scope_advice_targets( + target_cycles=10.0 if args.target_cycles is None else args.target_cycles, + target_vertical_divisions=( + 5.0 if args.target_vertical_divisions is None else args.target_vertical_divisions + ), + ) + channels = tuple(args.channels) if args.channels else None + expectations = _load_scope_expectations(args.expect) + if not args.fetch_waveform: + if expectations is not None: + raise ConfigError("scope observe --expect requires --fetch-waveform") + return scope_observe_payload( + config_path=args.config, + channels=channels, + allow_50ohm=args.allow_50ohm, + resource=args.resource, + ) + observation = scope_waveform_report_payload( + config_path=args.config, + channels=channels, + allow_50ohm=args.allow_50ohm, + expectations=expectations, + resource=args.resource, + ) + advice = scope_advise_from_observation( + observation, + expected_frequencies_hz=_expectation_frequencies(expectations), + target_cycles=target_cycles, + target_vertical_divisions=target_vertical_divisions, + ) + observation["recommendations"] = advice["recommendations"] + observation["agent_hints"] = advice["agent_hints"] + return observation + + +def _emit_scope_observe_result(payload: dict[str, Any], *, json_mode: bool) -> None: + if json_mode: + _emit_json_result(payload, status=str(payload.get("status", "ok"))) + return + print( + f"status={payload.get('status')} read_only={payload.get('read_only')} " + f"mutates_instrument={payload.get('mutates_instrument')}" + ) + identity = payload.get("identity") + if isinstance(identity, dict) and identity.get("status") == "ok": + print(f"idn={identity['data']['idn']}") + for channel_section in payload.get("channels", []) or []: + channel = channel_section.get("channel") + coupling = channel_section.get("coupling", {}) + coupling_value = coupling.get("data", {}).get("coupling") if coupling.get("status") == "ok" else "unavailable" + print(f"ch{channel} coupling={coupling_value}") + waveform = channel_section.get("waveform") + if isinstance(waveform, dict) and waveform.get("status") == "ok": + summary = waveform["data"]["summary"] + print( + f"ch{channel} waveform samples={summary.get('samples')} " + f"vpp_v={summary.get('voltage_vpp_v')} mean_v={summary.get('voltage_mean_v')} " + f"frequency_hz={summary.get('frequency_estimate_hz')}" + ) + for warning in summary.get("quality_warnings", []) or []: + print(f"ch{channel} quality_warning={warning}") + expectation = channel_section.get("expectation") + if isinstance(expectation, dict): + print(f"ch{channel} expectation={expectation.get('status')}") + for check in expectation.get("data", {}).get("checks", []) or []: + print( + f"ch{channel} check={check.get('metric')} status={check.get('status')} " + f"expected={check.get('expected')} actual={check.get('actual')}" + ) + waveform_source = payload.get("waveform_source") + if isinstance(waveform_source, dict) and not waveform_source.get("same_acquisition", True): + print(f"waveform_source same_acquisition=False reason={waveform_source.get('reason')}") + for relationship in payload.get("relationships", []) or []: + channels = relationship.get("channels") + frequency = relationship.get("frequency", {}) + print( + f"relationship ch{channels} ratio={frequency.get('ratio_high_over_low')} " + f"phase_deg={relationship.get('phase_degrees_at_left_frequency')}" + ) + for recommendation in payload.get("recommendations", []) or []: + command = recommendation.get("command") or recommendation.get("action") + print(f"recommendation {recommendation.get('id')} priority={recommendation.get('priority')} {command}") + for warning in payload.get("warnings", []) or []: + print(f"warning={warning}") + + def _scope_channel_display_request(args: argparse.Namespace) -> ScopeChannelDisplayRequest: try: return ScopeChannelDisplayRequest( @@ -1969,6 +2104,9 @@ def _main(argv: list[str] | None = None) -> int: print(f"summary={result.summary_path}") return 0 if args.domain == "scope": + if args.command == "observe": + _emit_scope_observe_result(_run_scope_observe(args), json_mode=args.json) + return 0 service = _load_service(args) if args.command == "idn": print(service.idn()) diff --git a/src/wavebench/cli_parser.py b/src/wavebench/cli_parser.py index f7c60b25..ab2d0f67 100644 --- a/src/wavebench/cli_parser.py +++ b/src/wavebench/cli_parser.py @@ -1546,6 +1546,46 @@ def add_trace_reference(parser: argparse.ArgumentParser) -> None: fetch.add_argument("--allow-50ohm", action="store_true", help="Explicitly allow scope input coupling that may be 50 ohm; default requires high impedance") add_runtime_options(fetch) + observe = scope_sub.add_parser( + "observe", + help="Observe scope state; --fetch-waveform explicitly reads waveforms for summaries and checks", + ) + observe.add_argument( + "--channel", + dest="channels", + type=int, + action="append", + default=None, + help="Observed analog channel; repeat for multiple channels", + ) + observe.add_argument( + "--fetch-waveform", + action="store_true", + help=( + "Explicitly read waveforms. This is a write path: it may stop a running acquisition, " + "change waveform transfer source/mode/format/points and enable channel display" + ), + ) + observe.add_argument( + "--allow-50ohm", + action="store_true", + help="Explicitly allow scope input coupling that may be 50 ohm; default requires high impedance", + ) + observe.add_argument( + "--expect", + default=None, + metavar="PATH", + help="TOML file with per-channel [channels.N] expectation checks; requires --fetch-waveform", + ) + observe.add_argument("--target-cycles", type=float, default=None, help="Target cycles for advice, default 10") + observe.add_argument( + "--target-vertical-divisions", + type=float, + default=None, + help="Target vertical divisions for advice, default 5", + ) + add_runtime_options(observe) + capture = scope_sub.add_parser("capture", help="Capture waveform data into an acquisition package") capture.add_argument("--channel", type=int, action="append", default=None, help="Capture channel; repeat for multiple channels") capture.add_argument("--label", default="capture") diff --git a/src/wavebench/data/expectations.py b/src/wavebench/data/expectations.py new file mode 100644 index 00000000..4ee836b5 --- /dev/null +++ b/src/wavebench/data/expectations.py @@ -0,0 +1,492 @@ +from __future__ import annotations + +import math +from typing import Any + +import numpy as np + +from wavebench.errors import ConfigError +from wavebench.instruments.models import WaveformData + +# expectation 允许的字符串字段 +_STRING_FIELDS = ("label", "shape") +# expectation 允许的数值字段:名称 -> (下界, 上界, 下界是否排他) +_NUMBER_FIELDS: dict[str, tuple[float, float, bool]] = { + "frequency_hz": (0.0, math.inf, True), + "frequency_tolerance_ratio": (0.0, math.inf, False), + "vpp_v": (0.0, math.inf, True), + "vpp_tolerance_ratio": (0.0, math.inf, False), + "mean_v": (-math.inf, math.inf, False), + "offset_v": (-math.inf, math.inf, False), + "mean_tolerance_v": (0.0, math.inf, False), + "duty_cycle": (0.0, 1.0, False), + "duty_percent": (0.0, 100.0, False), + "duty_tolerance": (0.0, math.inf, False), + "symmetry_percent": (0.0, 100.0, False), + "symmetry_tolerance_percent": (0.0, math.inf, False), +} +# 同义字段,不允许同时出现,避免“哪个生效”依赖字典顺序 +_MUTUALLY_EXCLUSIVE = (("duty_cycle", "duty_percent"), ("mean_v", "offset_v")) +# 三角波对称度估计的默认滞回阈值(占 Vpp 的比例) +_SYMMETRY_HYSTERESIS_RATIO = 0.02 + + +def validate_expectation(raw: dict[str, Any]) -> dict[str, Any]: + """严格校验 expectation,返回规范化副本;任何非法输入都抛出 ConfigError。 + + 这里必须自证,不能依赖 MCP 发布的 JSON Schema:schema 只在客户端生效。 + """ + if not isinstance(raw, dict): + raise ConfigError("expectation entries must be objects / expectation 条目必须是对象") + known = set(_STRING_FIELDS) | set(_NUMBER_FIELDS) + unknown = sorted(set(raw) - known) + if unknown: + raise ConfigError( + f"unknown expectation field(s): {', '.join(unknown)} / expectation 存在未知字段" + ) + validated: dict[str, Any] = {} + for name in _STRING_FIELDS: + value = raw.get(name) + if value is None: + continue + if not isinstance(value, str) or not value.strip(): + raise ConfigError(f"expectation {name} must be a non-empty string") + validated[name] = value + for name, (low, high, exclusive_low) in _NUMBER_FIELDS.items(): + value = raw.get(name) + if value is None: + continue + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ConfigError(f"expectation {name} must be a number") + number = float(value) + if not math.isfinite(number): + raise ConfigError(f"expectation {name} must be finite") + if exclusive_low and number <= low: + raise ConfigError(f"expectation {name} must be > {low:g}") + if not exclusive_low and number < low: + raise ConfigError(f"expectation {name} must be >= {low:g}") + if number > high: + raise ConfigError(f"expectation {name} must be <= {high:g}") + validated[name] = number + for left, right in _MUTUALLY_EXCLUSIVE: + if left in validated and right in validated: + raise ConfigError(f"expectation must not set both {left} and {right}") + return validated + + +def evaluate_waveform_expectation( + waveform: WaveformData, + expectation: dict[str, Any], +) -> dict[str, Any]: + validated = validate_expectation(expectation) + summary = waveform.summary( + expected_frequency_hz=_number(validated, "frequency_hz"), + frequency_tolerance_ratio=_number(validated, "frequency_tolerance_ratio", default=0.05), + ) + checks: list[dict[str, Any]] = [] + _check_frequency(summary, validated, checks) + _check_vpp(summary, validated, checks) + _check_mean(summary, validated, checks) + _check_duty(summary, validated, checks) + _check_symmetry(waveform, validated, checks) + if not checks: + # 没有任何可执行检查时必须显式说明,不能伪装成 pass + return { + "status": "skipped", + "channel": waveform.channel, + "label": validated.get("label"), + "shape": validated.get("shape"), + "checks": [], + "message": "expectation contains no checkable metric", + } + statuses = {check["status"] for check in checks} + if "fail" in statuses: + status = "fail" + elif "warn" in statuses: + status = "warn" + else: + status = "pass" + return { + "status": status, + "channel": waveform.channel, + "label": validated.get("label"), + "shape": validated.get("shape"), + "checks": checks, + } + + +def expectation_summary(results: dict[int, dict[str, Any]]) -> dict[str, Any]: + """Keep confirmed failures; incomplete acceptance is never a pass.""" + statuses = {result["status"] for result in results.values()} + if "fail" in statuses: + status = "fail" + elif "unavailable" in statuses: + status = "partial" if statuses & {"pass", "warn"} else "unavailable" + elif "warn" in statuses: + status = "warn" + elif "pass" in statuses: + status = "pass" + else: + status = "skipped" + return { + "status": status, + "channels": {str(channel): result["status"] for channel, result in sorted(results.items())}, + } + + +def estimate_triangle_symmetry_percent( + waveform: WaveformData, + *, + expected_frequency_hz: float | None = None, + hysteresis_ratio: float = _SYMMETRY_HYSTERESIS_RATIO, +) -> float | None: + """估计三角波上升沿占整周期的百分比。 + + 真实示波器波形带有噪声、量化台阶和过冲,逐点比较相邻差分符号会被单个噪声样本打乱。 + 这里先用滑动平均抑制噪声,再用滞回(相对 Vpp)确认极值反转,并用期望频率约束周期长度。 + """ + times = waveform.times_s + values = np.asarray(waveform.voltages_v, dtype=np.float64) + if times.size != values.size or values.size < 8: + return None + span = float(np.max(values) - np.min(values)) + if span <= 1e-12: + return None + samples_per_cycle = _samples_per_cycle(times, expected_frequency_hz) + smoothed = _smooth(values, _smoothing_window(values.size, samples_per_cycle)) + extrema = _hysteresis_extrema(smoothed, hysteresis=max(span * hysteresis_ratio, 0.0)) + if len(extrema) < 3: + return None + # 滞回和滑窗都会把极值确认点推向信号内部;用相邻两段原始数据的拟合直线交点把 + # 极值时间还原回三角波的真实折点。 + refined = _refine_extrema(times, values, extrema) + min_period_s = 0.5 / expected_frequency_hz if expected_frequency_hz else None + fractions: list[float] = [] + for index in range(1, len(refined) - 1): + left_kind, left_time = refined[index - 1] + kind, peak_time = refined[index] + right_kind, right_time = refined[index + 1] + if (left_kind, kind, right_kind) != ("min", "max", "min"): + continue + period = right_time - left_time + if period <= 0 or (min_period_s is not None and period < min_period_s): + continue + fractions.append(float((peak_time - left_time) / period * 100.0)) + if not fractions: + return None + return float(np.median(np.asarray(fractions, dtype=np.float64))) + + +def _refine_extrema( + times: np.ndarray, + values: np.ndarray, + extrema: list[tuple[str, int]], +) -> list[tuple[str, float]]: + refined: list[tuple[str, float]] = [] + for position, (kind, index) in enumerate(extrema): + if position == 0 or position == len(extrema) - 1: + refined.append((kind, float(times[index]))) + continue + crossing = _segment_intersection_time( + times, + values, + extrema[position - 1][1], + index, + extrema[position + 1][1], + ) + refined.append((kind, crossing if crossing is not None else float(times[index]))) + return refined + + +def _segment_intersection_time( + times: np.ndarray, + values: np.ndarray, + previous_index: int, + index: int, + next_index: int, +) -> float | None: + rising = _fit_line(times, values, previous_index, index) + falling = _fit_line(times, values, index, next_index) + if rising is None or falling is None: + return None + slope_a, intercept_a = rising + slope_b, intercept_b = falling + if abs(slope_a - slope_b) <= 1e-18: + return None + crossing = (intercept_b - intercept_a) / (slope_a - slope_b) + if not math.isfinite(crossing): + return None + return float(crossing) + + +def _fit_line( + times: np.ndarray, + values: np.ndarray, + start_index: int, + stop_index: int, +) -> tuple[float, float] | None: + low, high = (start_index, stop_index) if start_index <= stop_index else (stop_index, start_index) + if high - low < 2: + return None + # 裁掉两端靠近折点的部分,避免拐角处的采样点把拟合斜率拽偏 + margin = max(1, int((high - low) * 0.15)) + begin = low + margin + end = high - margin + if end - begin < 1: + begin, end = low, high + x = np.asarray(times[begin : end + 1], dtype=np.float64) + y = np.asarray(values[begin : end + 1], dtype=np.float64) + if x.size < 2: + return None + slope, intercept = np.polyfit(x, y, 1) + return float(slope), float(intercept) + + +def _samples_per_cycle(times: np.ndarray, expected_frequency_hz: float | None) -> float | None: + if not expected_frequency_hz or expected_frequency_hz <= 0 or times.size < 2: + return None + duration = float(times[-1] - times[0]) + if duration <= 0: + return None + return float(times.size) / (duration * expected_frequency_hz) + + +def _smoothing_window(size: int, samples_per_cycle: float | None) -> int: + if samples_per_cycle is not None and samples_per_cycle > 64: + candidate = int(samples_per_cycle // 20) + else: + candidate = int(size // 50) + window = max(1, min(candidate, 51, size)) + if window > 1 and window % 2 == 0: + window -= 1 + return window + + +def _smooth(values: np.ndarray, window: int) -> np.ndarray: + if window <= 1: + return values + kernel = np.ones(window, dtype=np.float64) / window + pad = window // 2 + padded = np.pad(values, pad, mode="edge") + return np.convolve(padded, kernel, mode="valid") + + +def _hysteresis_extrema(values: np.ndarray, *, hysteresis: float) -> list[tuple[str, int]]: + """用滞回跟踪方向变化,返回 [(kind, index), ...],kind 为 min/max 且交替出现。 + + 方向未确认时锚点保持不变,否则单调斜坡会被“锚点跟着当前值走”抵消掉滞回判据。 + """ + extrema: list[tuple[str, int]] = [] + if values.size == 0: + return extrema + anchor = 0 + direction = 0 + for index in range(1, values.size): + value = float(values[index]) + if direction == 0: + if value >= float(values[anchor]) + hysteresis: + direction = 1 + elif value <= float(values[anchor]) - hysteresis: + direction = -1 + continue + if direction > 0: + if value > float(values[anchor]): + anchor = index + elif float(values[anchor]) - value > hysteresis: + if anchor > 0: + extrema.append(("max", anchor)) + direction, anchor = -1, index + else: + if value < float(values[anchor]): + anchor = index + elif value - float(values[anchor]) > hysteresis: + if anchor > 0: + extrema.append(("min", anchor)) + direction, anchor = 1, index + return extrema + + +def _check_frequency( + summary: dict[str, Any], + expectation: dict[str, Any], + checks: list[dict[str, Any]], +) -> None: + expected = _number(expectation, "frequency_hz") + if expected is None: + return + actual = summary.get("frequency_estimate_hz") + tolerance = _number(expectation, "frequency_tolerance_ratio", default=0.05) + low_confidence = any( + str(item).startswith("low_cycle_count") + for item in summary.get("quality_warnings", []) + ) + if not isinstance(actual, (int, float)) or actual <= 0: + checks.append(_check("frequency_hz", "warn", expected, actual, "frequency unavailable")) + return + error_ratio = abs(float(actual) - expected) / expected + if low_confidence: + checks.append( + _check( + "frequency_hz", + "warn", + expected, + float(actual), + "frequency low confidence because waveform contains too few cycles", + error_ratio=error_ratio, + tolerance_ratio=tolerance, + ) + ) + return + checks.append( + _check( + "frequency_hz", + "pass" if error_ratio <= tolerance else "fail", + expected, + float(actual), + "ok" if error_ratio <= tolerance else "frequency out of tolerance", + error_ratio=error_ratio, + tolerance_ratio=tolerance, + ) + ) + + +def _check_vpp( + summary: dict[str, Any], + expectation: dict[str, Any], + checks: list[dict[str, Any]], +) -> None: + expected = _number(expectation, "vpp_v") + if expected is None: + return + actual = summary.get("voltage_vpp_v") + tolerance = _number(expectation, "vpp_tolerance_ratio", default=0.10) + if not isinstance(actual, (int, float)): + checks.append(_check("vpp_v", "warn", expected, actual, "Vpp unavailable")) + return + error_ratio = abs(float(actual) - expected) / expected + checks.append( + _check( + "vpp_v", + "pass" if error_ratio <= tolerance else "fail", + expected, + float(actual), + "ok" if error_ratio <= tolerance else "Vpp out of tolerance", + error_ratio=error_ratio, + tolerance_ratio=tolerance, + ) + ) + + +def _check_mean( + summary: dict[str, Any], + expectation: dict[str, Any], + checks: list[dict[str, Any]], +) -> None: + expected = _number(expectation, "mean_v", fallback_field="offset_v") + if expected is None: + return + actual = summary.get("voltage_mean_v") + tolerance = _number(expectation, "mean_tolerance_v", default=0.05) + if not isinstance(actual, (int, float)): + checks.append(_check("mean_v", "warn", expected, actual, "mean unavailable")) + return + error = abs(float(actual) - expected) + checks.append( + _check( + "mean_v", + "pass" if error <= tolerance else "fail", + expected, + float(actual), + "ok" if error <= tolerance else "mean out of tolerance", + error_abs=error, + tolerance_abs=tolerance, + ) + ) + + +def _check_duty( + summary: dict[str, Any], + expectation: dict[str, Any], + checks: list[dict[str, Any]], +) -> None: + expected = _number(expectation, "duty_cycle") + if expected is None: + percent = _number(expectation, "duty_percent") + if percent is not None: + expected = percent / 100.0 + if expected is None: + return + actual = summary.get("duty_cycle") + tolerance = _number(expectation, "duty_tolerance", default=0.05) + if not isinstance(actual, (int, float)): + checks.append(_check("duty_cycle", "warn", expected, actual, "duty unavailable")) + return + error = abs(float(actual) - expected) + checks.append( + _check( + "duty_cycle", + "pass" if error <= tolerance else "fail", + expected, + float(actual), + "ok" if error <= tolerance else "duty out of tolerance", + error_abs=error, + tolerance_abs=tolerance, + ) + ) + + +def _check_symmetry( + waveform: WaveformData, + expectation: dict[str, Any], + checks: list[dict[str, Any]], +) -> None: + expected = _number(expectation, "symmetry_percent") + if expected is None: + return + actual = estimate_triangle_symmetry_percent( + waveform, + expected_frequency_hz=_number(expectation, "frequency_hz"), + ) + tolerance = _number(expectation, "symmetry_tolerance_percent", default=5.0) + if actual is None: + checks.append(_check("symmetry_percent", "warn", expected, actual, "symmetry unavailable")) + return + error = abs(actual - expected) + checks.append( + _check( + "symmetry_percent", + "pass" if error <= tolerance else "fail", + expected, + actual, + "ok" if error <= tolerance else "symmetry out of tolerance", + error_abs=error, + tolerance_abs=tolerance, + ) + ) + + +def _check(name: str, status: str, expected: Any, actual: Any, message: str, **extra: Any) -> dict[str, Any]: + return { + "metric": name, + "status": status, + "expected": expected, + "actual": actual, + "message": message, + **extra, + } + + +def _number( + expectation: dict[str, Any], + name: str, + *, + default: float | None = None, + fallback_field: str | None = None, +) -> float | None: + """读取已校验的数值字段。``fallback_field`` 只用于同义字段,两者不会同时出现。""" + value = expectation.get(name) + if value is None and fallback_field is not None: + value = expectation.get(fallback_field) + if value is None: + return default + return float(value) diff --git a/src/wavebench/data/relationships.py b/src/wavebench/data/relationships.py new file mode 100644 index 00000000..943a92cb --- /dev/null +++ b/src/wavebench/data/relationships.py @@ -0,0 +1,344 @@ +from __future__ import annotations + +import math +from itertools import combinations +from typing import Any + +import numpy as np + +from wavebench.instruments.models import WaveformData + + +def analyze_waveform_relationships( + waveforms: dict[int, WaveformData], + *, + same_acquisition: bool = True, + max_correlation_points: int = 4096, + max_intersections: int = 64, +) -> list[dict[str, Any]]: + relationships: list[dict[str, Any]] = [] + for left_channel, right_channel in combinations(sorted(waveforms), 2): + relationships.append( + analyze_waveform_pair( + waveforms[left_channel], + waveforms[right_channel], + same_acquisition=same_acquisition, + max_correlation_points=max_correlation_points, + max_intersections=max_intersections, + ) + ) + return relationships + + +def analyze_waveform_pair( + left: WaveformData, + right: WaveformData, + *, + same_acquisition: bool = True, + max_correlation_points: int = 4096, + max_intersections: int = 64, +) -> dict[str, Any]: + left_summary = left.summary() + right_summary = right.summary() + warnings: list[str] = [] + if same_acquisition: + common = _common_time_axis(left, right, max_points=max_correlation_points) + correlation = _correlation_payload(common, warnings=warnings) + intersections = _intersection_payload( + common, + warnings=warnings, + max_intersections=max_intersections, + ) + common_time = {**common["metadata"], "same_acquisition": True} + else: + # 跨 acquisition 的两个通道没有共同时间基准:相关性、交点、相位、延迟都不成立, + # 只看同步无关的频率比和幅度/均值关系。 + warnings.append("not_same_acquisition_timing_relationships_skipped") + correlation = _skipped_analysis("not_same_acquisition") + intersections = _skipped_analysis("not_same_acquisition") + common_time = { + "overlap": None, + "x_start_s": None, + "x_stop_s": None, + "duration_s": None, + "samples": 0, + "same_acquisition": False, + } + left_frequency = _trusted_frequency(left_summary, warnings=warnings, label=f"CH{left.channel}") + right_frequency = _trusted_frequency(right_summary, warnings=warnings, label=f"CH{right.channel}") + frequency_ratio = None + phase_degrees = None + if left_frequency is not None and right_frequency is not None: + lower = min(left_frequency, right_frequency) + upper = max(left_frequency, right_frequency) + if lower > 0: + frequency_ratio = float(upper / lower) + if not same_acquisition: + pass + elif ( + abs(left_frequency - right_frequency) / max(left_frequency, right_frequency) <= 0.01 + and common_time.get("overlap") is True + ): + # 约定:phase_degrees_at_left_frequency 表示 right 相对 left 的相位滞后,取值 [0, 360)。 + # 用基波拟合相位差而不是相关峰 lag:后者对截断窗口和幅度不对称有系统偏差, + # 且直接反相(right = -left)会被绝对值最大化吃掉 180°。 + phase_degrees = _fundamental_phase_degrees(common, frequency_hz=left_frequency) + elif frequency_ratio is not None and abs(frequency_ratio - 1.0) > 0.01: + warnings.append("phase_not_meaningful_for_different_frequencies") + return { + "channels": [left.channel, right.channel], + "left_channel": left.channel, + "right_channel": right.channel, + "common_time": common_time, + "frequency": { + "left_hz": left_frequency, + "right_hz": right_frequency, + "ratio_high_over_low": frequency_ratio, + }, + "voltage": { + "left_vpp_v": left_summary["voltage_vpp_v"], + "right_vpp_v": right_summary["voltage_vpp_v"], + "vpp_ratio_right_over_left": _safe_ratio( + right_summary["voltage_vpp_v"], + left_summary["voltage_vpp_v"], + ), + "mean_delta_right_minus_left_v": float( + right_summary["voltage_mean_v"] - left_summary["voltage_mean_v"] + ), + "rms_ratio_right_over_left": _safe_ratio( + right_summary["voltage_rms_v"], + left_summary["voltage_rms_v"], + ), + }, + "correlation": correlation, + "intersections": intersections, + "phase_degrees_at_left_frequency": phase_degrees, + "warnings": warnings, + } + + +def _common_time_axis( + left: WaveformData, + right: WaveformData, + *, + max_points: int, +) -> dict[str, Any]: + left_times = left.times_s + right_times = right.times_s + start = max(float(left_times[0]), float(right_times[0])) + stop = min(float(left_times[-1]), float(right_times[-1])) + if stop <= start: + return { + "time_s": np.array([], dtype=np.float64), + "left_v": np.array([], dtype=np.float64), + "right_v": np.array([], dtype=np.float64), + "metadata": { + "overlap": False, + "x_start_s": start, + "x_stop_s": stop, + "duration_s": 0.0, + "samples": 0, + }, + } + left_dt = left.header.x_increment + right_dt = right.header.x_increment + dt = max(value for value in (left_dt, right_dt) if value > 0) + count = int(np.floor((stop - start) / dt)) + 1 + count = max(2, min(count, max_points)) + common_times = np.linspace(start, stop, count, dtype=np.float64) + return { + "time_s": common_times, + "left_v": np.interp(common_times, left_times, left.voltages_v), + "right_v": np.interp(common_times, right_times, right.voltages_v), + "metadata": { + "overlap": True, + "x_start_s": start, + "x_stop_s": stop, + "duration_s": float(stop - start), + "samples": count, + }, + } + + +def _skipped_analysis(reason: str) -> dict[str, Any]: + return {"status": "skipped", "reason": reason} + + +def _fundamental_phase_degrees(common: dict[str, Any], *, frequency_hz: float) -> float | None: + """在 common_time 上拟合基波,返回 right 相对 left 的相位滞后(度,[0, 360))。""" + times = common["time_s"] + left = np.asarray(common["left_v"], dtype=np.float64) + right = np.asarray(common["right_v"], dtype=np.float64) + if frequency_hz <= 0 or times.size < 8: + return None + phase_left = _single_bin_phase(times, left, frequency_hz) + phase_right = _single_bin_phase(times, right, frequency_hz) + if phase_left is None or phase_right is None: + return None + return float((-math.degrees(_wrap_angle(phase_right - phase_left))) % 360.0) + + +def _single_bin_phase(times: np.ndarray, values: np.ndarray, frequency_hz: float) -> float | None: + # 非整数周期窗口内常数、cos、sin 不正交,必须联合拟合以消除 DC 泄漏。 + angle = 2.0 * math.pi * frequency_hz * (times - times[0]) + basis = np.column_stack((np.ones_like(angle), np.cos(angle), np.sin(angle))) + coefficients, _, rank, _ = np.linalg.lstsq(basis, values, rcond=None) + _, cosine, sine = coefficients + tolerance = np.finfo(np.float64).eps * max(float(np.max(np.abs(values))), 1.0) * 8 + if rank < 3 or math.hypot(cosine, sine) <= tolerance: + return None + return math.atan2(-sine, cosine) + + +def _wrap_angle(angle: float) -> float: + wrapped = (angle + math.pi) % (2.0 * math.pi) + return wrapped - math.pi + + +def _correlation_payload(common: dict[str, Any], *, warnings: list[str]) -> dict[str, Any]: + times = common["time_s"] + if times.size < 4: + warnings.append("insufficient_common_time_overlap") + return { + "normalized_pearson": None, + "max_cross_correlation": None, + "max_abs_cross_correlation": None, + "lag_at_max_correlation_s": None, + } + left = _normalize(common["left_v"]) + right = _normalize(common["right_v"]) + if left is None or right is None: + warnings.append("correlation_unavailable_for_flat_signal") + return { + "normalized_pearson": None, + "max_cross_correlation": None, + "max_abs_cross_correlation": None, + "lag_at_max_correlation_s": None, + } + pearson = float(np.mean(left * right)) + correlation = np.correlate(right, left, mode="full") / left.size + index = int(np.argmax(np.abs(correlation))) + lag_samples = index - (left.size - 1) + dt = float(np.median(np.diff(times))) + return { + "normalized_pearson": pearson, + "max_cross_correlation": float(correlation[index]), + "max_abs_cross_correlation": float(abs(correlation[index])), + "lag_at_max_correlation_s": float(lag_samples * dt), + } + + +def _intersection_payload( + common: dict[str, Any], + *, + warnings: list[str], + max_intersections: int, +) -> dict[str, Any]: + times = common["time_s"] + left = common["left_v"] + right = common["right_v"] + if times.size < 2: + return { + "mode": "none", + "count": 0, + "returned": 0, + "truncated": False, + "points": [], + } + diff = left - right + tolerance = max(float(np.max(np.abs(diff))) * 1e-9, 1e-12) + if bool(np.all(np.abs(diff) <= tolerance)): + warnings.append("waveforms_coincident_intersections_unbounded") + return { + "mode": "coincident", + "count": None, + "returned": 0, + "truncated": False, + "points": [], + } + points: list[dict[str, float | str]] = [] + count = 0 + last_time: float | None = None + for index in range(diff.size - 1): + d0 = float(diff[index]) + d1 = float(diff[index + 1]) + t0 = float(times[index]) + t1 = float(times[index + 1]) + if abs(d0) <= tolerance: + alpha = 0.0 + elif d0 * d1 < 0.0: + alpha = -d0 / (d1 - d0) + else: + continue + crossing_time = t0 + alpha * (t1 - t0) + if last_time is not None and abs(crossing_time - last_time) <= max(abs(t1 - t0) * 0.5, 1e-15): + continue + left_value = float(left[index] + alpha * (left[index + 1] - left[index])) + right_value = float(right[index] + alpha * (right[index + 1] - right[index])) + left_slope = _segment_slope(left, times, index) + right_slope = _segment_slope(right, times, index) + delta_slope = left_slope - right_slope + count += 1 + last_time = crossing_time + if len(points) < max_intersections: + points.append( + { + "time_s": float(crossing_time), + "voltage_v": float((left_value + right_value) / 2.0), + "left_slope_v_per_s": float(left_slope), + "right_slope_v_per_s": float(right_slope), + "delta_slope_v_per_s": float(delta_slope), + "direction": ( + "left_minus_right_rising" + if delta_slope > 0 + else "left_minus_right_falling" + if delta_slope < 0 + else "tangent_or_flat" + ), + } + ) + truncated = count > len(points) + if truncated: + warnings.append("intersections_truncated") + return { + "mode": "finite", + "count": count, + "returned": len(points), + "truncated": truncated, + "points": points, + } + + +def _segment_slope(values: np.ndarray, times: np.ndarray, index: int) -> float: + dt = float(times[index + 1] - times[index]) + if abs(dt) <= 1e-18: + return 0.0 + return float((values[index + 1] - values[index]) / dt) + + +def _normalize(values: np.ndarray) -> np.ndarray | None: + centered = values.astype(np.float64) - float(np.mean(values)) + rms = float(np.sqrt(np.mean(np.square(centered)))) + if rms <= 1e-12: + return None + return centered / rms + + +def _trusted_frequency(summary: dict[str, object], *, warnings: list[str], label: str) -> float | None: + frequency = summary.get("frequency_estimate_hz") + if not isinstance(frequency, (int, float)) or frequency <= 0: + warnings.append(f"{label}_frequency_unavailable") + return None + quality_warnings = summary.get("quality_warnings", []) + if any(str(item).startswith("low_cycle_count") for item in quality_warnings): + warnings.append(f"{label}_frequency_low_confidence") + return None + return float(frequency) + + +def _safe_ratio(numerator: object, denominator: object) -> float | None: + if not isinstance(numerator, (int, float)) or not isinstance(denominator, (int, float)): + return None + if abs(float(denominator)) <= 1e-18: + return None + return float(numerator) / float(denominator) diff --git a/src/wavebench/mcp_http.py b/src/wavebench/mcp_http.py index 4f3a344a..7683eaf8 100644 --- a/src/wavebench/mcp_http.py +++ b/src/wavebench/mcp_http.py @@ -1,6 +1,7 @@ from __future__ import annotations import json +import math import os from dataclasses import dataclass from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer @@ -12,8 +13,11 @@ from wavebench.config import load_config from wavebench import __version__ from wavebench.data.packages import load_capture_package +from wavebench.doctor import doctor_records, has_doctor_errors from wavebench.errors import ConfigError, WaveBenchError from wavebench.logging import CommandLogger +from wavebench.services.agent_advise import scope_advise_payload +from wavebench.services.agent_observe import scope_observe_payload from wavebench.services.run_plan import load_run_plan, run_plan_schema_rows from wavebench.services.run_plan import format_run_plan_schema from wavebench.services.run_service import RunService @@ -50,12 +54,19 @@ class ToolSpec: description: str arguments: dict[str, Any] handler: Callable[[dict[str, Any], Path], dict[str, Any]] + read_only: bool = True + mutates_instrument: bool = False + raw_scpi: bool = False + instrument_state_effects: tuple[str, ...] = () def public_payload(self) -> dict[str, Any]: return { "name": self.name, "description": self.description, - "read_only": True, + "read_only": self.read_only, + "mutates_instrument": self.mutates_instrument, + "raw_scpi": self.raw_scpi, + "instrument_state_effects": list(self.instrument_state_effects), "arguments": self.arguments, } @@ -216,6 +227,119 @@ def _capture_inspect_tool(arguments: dict[str, Any], config_path: Path) -> dict[ } +def _optional_bool(arguments: dict[str, Any], name: str, default: bool) -> bool: + value = arguments.get(name, default) + if not isinstance(value, bool): + raise ConfigError(f"{name} must be a boolean / {name} 必须是布尔值") + return value + + +def _scope_observe_tool(arguments: dict[str, Any], config_path: Path) -> dict[str, Any]: + channel, channels = _scope_channel_arguments(arguments) + return scope_observe_payload( + config_path=_reject_sensitive_path(config_path, label="config"), + channel=channel, + channels=channels, + allow_50ohm=_optional_bool(arguments, "allow_50ohm", False), + ) + + +def _scope_advise_tool(arguments: dict[str, Any], config_path: Path) -> dict[str, Any]: + channel, channels = _scope_channel_arguments(arguments) + return scope_advise_payload( + config_path=_reject_sensitive_path(config_path, label="config"), + channel=channel, + channels=channels, + allow_50ohm=_optional_bool(arguments, "allow_50ohm", False), + expected_frequencies_hz=_expected_frequencies_argument( + arguments.get("expected_frequencies_hz") + ), + target_cycles=_optional_positive_number(arguments, "target_cycles", 10.0), + target_vertical_divisions=_optional_positive_number( + arguments, + "target_vertical_divisions", + 5.0, + ), + ) + + +def _scope_channel_arguments(arguments: dict[str, Any]) -> tuple[int | None, tuple[int, ...] | None]: + channel = arguments.get("channel") + if channel is not None and (isinstance(channel, bool) or not isinstance(channel, int)): + raise ConfigError("channel must be an integer / channel 必须是整数") + raw_channels = arguments.get("channels") + channels = None + if raw_channels is not None: + if not isinstance(raw_channels, list): + raise ConfigError("channels must be an array / channels 必须是数组") + channels = tuple(raw_channels) + if any(isinstance(item, bool) or not isinstance(item, int) for item in channels): + raise ConfigError("channels must contain integers / channels 必须包含整数") + return channel, channels + + +def _optional_positive_number(arguments: dict[str, Any], name: str, default: float) -> float: + value = arguments.get(name, default) + if ( + isinstance(value, bool) + or not isinstance(value, (int, float)) + or not math.isfinite(float(value)) + or value <= 0 + ): + raise ConfigError(f"{name} must be a finite positive number / {name} 必须是有限正数") + return float(value) + + +def _expected_frequencies_argument(raw: Any) -> dict[int, float] | None: + if raw is None: + return None + if not isinstance(raw, dict): + raise ConfigError("expected_frequencies_hz must be an object / expected_frequencies_hz 必须是对象") + parsed: dict[int, float] = {} + for key, value in raw.items(): + try: + channel = int(key) + except (TypeError, ValueError) as exc: + raise ConfigError( + "expected_frequencies_hz keys must be channel numbers / expected_frequencies_hz 键必须是通道号" + ) from exc + if channel < 1: + raise ConfigError("expected_frequencies_hz channel must be >= 1 / 通道必须 >= 1") + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ConfigError("expected_frequencies_hz values must be numbers / 频率必须是数字") + parsed[channel] = float(value) + return parsed + + +def _doctor_config_tool(arguments: dict[str, Any], config_path: Path) -> dict[str, Any]: + timeout_ms = arguments.get("timeout_ms") + if timeout_ms is not None and (isinstance(timeout_ms, bool) or not isinstance(timeout_ms, int) or timeout_ms <= 0): + raise ConfigError("timeout_ms must be a positive integer / timeout_ms 必须是正整数") + records = doctor_records( + load_config(_reject_sensitive_path(config_path, label="config")), + timeout_ms=timeout_ms, + include_visa=False, + ) + return { + "status": "error" if has_doctor_errors(records) else "ok", + "read_only": True, + "mutates_instrument": False, + "raw_scpi": False, + "records": [ + { + "severity": record.severity, + "target": record.target, + "driver": record.driver, + "resource": record.resource, + "idn": record.idn, + "message": record.message, + "suggestion": record.suggestion, + } + for record in records + ], + } + + READ_ONLY_TOOLS: dict[str, ToolSpec] = { "run.schema": ToolSpec( name="run.schema", @@ -248,6 +372,78 @@ def _capture_inspect_tool(arguments: dict[str, Any], config_path: Path) -> dict[ }, handler=_capture_inspect_tool, ), + "scope.observe": ToolSpec( + name="scope.observe", + description=( + "Read configured scope identity, per-channel state snapshot, and input-coupling safety. " + "Never reads waveforms and never changes instrument state / " + "只读观察配置中的示波器身份、通道状态与输入耦合安全;不读取波形,不改变仪器状态" + ), + arguments={ + "type": "object", + "properties": { + "channel": {"type": "integer", "minimum": 1}, + "channels": { + "type": "array", + "items": {"type": "integer", "minimum": 1}, + "minItems": 1, + "uniqueItems": True, + }, + "allow_50ohm": {"type": "boolean", "default": False}, + }, + "additionalProperties": False, + }, + handler=_scope_observe_tool, + ), + "doctor.config": ToolSpec( + name="doctor.config", + description=( + "Run configured-instrument read-only doctor checks and return structured records / " + "对已配置仪器执行只读 doctor 检查并返回结构化结果" + ), + arguments={ + "type": "object", + "properties": { + "timeout_ms": {"type": "integer", "minimum": 1}, + }, + "additionalProperties": False, + }, + handler=_doctor_config_tool, + ), + "scope.advise": ToolSpec( + name="scope.advise", + description=( + "Observe the configured scope and recommend display/acquisition settings from the " + "read-only state snapshot and caller-provided expected frequencies; never reads waveforms " + "and never applies recommendations / " + "基于只读状态快照与调用方给出的期望频率建议显示/采集参数;不读取波形,不应用建议" + ), + arguments={ + "type": "object", + "properties": { + "channel": {"type": "integer", "minimum": 1}, + "channels": { + "type": "array", + "items": {"type": "integer", "minimum": 1}, + "minItems": 1, + "uniqueItems": True, + }, + "allow_50ohm": {"type": "boolean", "default": False}, + "target_cycles": {"type": "number", "exclusiveMinimum": 0, "default": 10}, + "target_vertical_divisions": { + "type": "number", + "exclusiveMinimum": 0, + "default": 5, + }, + "expected_frequencies_hz": { + "type": "object", + "additionalProperties": {"type": "number", "exclusiveMinimum": 0}, + }, + }, + "additionalProperties": False, + }, + handler=_scope_advise_tool, + ), } diff --git a/src/wavebench/services/agent_advise.py b/src/wavebench/services/agent_advise.py new file mode 100644 index 00000000..176c2467 --- /dev/null +++ b/src/wavebench/services/agent_advise.py @@ -0,0 +1,434 @@ +from __future__ import annotations + +import math +from pathlib import Path +from typing import Any + +from wavebench.errors import ConfigError +from wavebench.services.agent_observe import scope_observe_payload + + +def scope_advise_payload( + *, + config_path: str | Path, + channel: int | None = None, + channels: tuple[int, ...] | None = None, + allow_50ohm: bool = False, + expected_frequencies_hz: dict[int, float] | None = None, + target_cycles: float = 10.0, + target_vertical_divisions: float = 5.0, + resource: str | None = None, +) -> dict[str, Any]: + """只读建议:基于示波器状态快照和调用方提供的期望频率给出显示/时基建议。 + + 该函数不读取波形、不改变仪器状态,也不基于低置信度的测量频率下结论。 + 需要基于实测波形的建议时,使用 ``scope_waveform_report_payload`` 的结果调用 + ``scope_advise_from_observation``。 + """ + # 参数校验必须发生在打开任何仪器会话之前 + target_cycles, target_vertical_divisions = validate_scope_advice_targets( + target_cycles=target_cycles, + target_vertical_divisions=target_vertical_divisions, + ) + expected_frequencies = _normalize_expected_frequencies(expected_frequencies_hz) + observation = scope_observe_payload( + config_path=config_path, + channel=channel, + channels=channels, + allow_50ohm=allow_50ohm, + resource=resource, + ) + return scope_advise_from_observation( + observation, + expected_frequencies_hz=expected_frequencies, + target_cycles=target_cycles, + target_vertical_divisions=target_vertical_divisions, + ) + + +def scope_advise_from_observation( + observation: dict[str, Any], + *, + expected_frequencies_hz: dict[int, float] | None = None, + target_cycles: float = 10.0, + target_vertical_divisions: float = 5.0, +) -> dict[str, Any]: + target_cycles, target_vertical_divisions = validate_scope_advice_targets( + target_cycles=target_cycles, + target_vertical_divisions=target_vertical_divisions, + ) + expected = _normalize_expected_frequencies(expected_frequencies_hz) + recommendations = _recommendations( + observation, + expected_frequencies=expected, + target_cycles=target_cycles, + target_vertical_divisions=target_vertical_divisions, + ) + return { + "status": observation["status"], + "read_only": observation["read_only"], + "query_only": observation["query_only"], + "mutates_instrument": observation["mutates_instrument"], + "raw_scpi": False, + "applies_recommendations": False, + "instrument_state_effects": observation["instrument_state_effects"], + "observation": { + "channel": observation["observation"]["channel"], + "channels": observation["observation"]["channels"], + "fetch_waveform": observation["observation"]["fetch_waveform"], + }, + "expected_frequencies_hz": {str(item): value for item, value in sorted(expected.items())}, + "recommendations": recommendations, + "agent_hints": _agent_hints(observation, recommendations), + "warnings": observation["warnings"], + } + + +def validate_scope_advice_targets( + *, target_cycles: float, target_vertical_divisions: float, +) -> tuple[float, float]: + """Validate advice targets before observation or instrument I/O.""" + return ( + _positive_finite(target_cycles, name="scope.advise target_cycles"), + _positive_finite(target_vertical_divisions, name="scope.advise target_vertical_divisions"), + ) + + +def _normalize_expected_frequencies(values: dict[int, float] | None) -> dict[int, float]: + if values is None: + return {} + normalized: dict[int, float] = {} + for channel, value in values.items(): + if isinstance(channel, bool) or not isinstance(channel, int) or channel < 1: + raise ConfigError("expected frequency channel must be a positive integer") + normalized[channel] = _positive_finite(value, name=f"expected frequency for channel {channel}") + return normalized + + +def _positive_finite(value: Any, *, name: str) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ConfigError(f"{name} must be a number") + number = float(value) + if not math.isfinite(number) or number <= 0: + raise ConfigError(f"{name} must be finite and > 0") + return number + + +def _recommendations( + observation: dict[str, Any], + *, + expected_frequencies: dict[int, float], + target_cycles: float, + target_vertical_divisions: float, +) -> list[dict[str, Any]]: + recommendations: list[dict[str, Any]] = [] + channels = observation.get("channels", []) + channel_profiles: dict[int, dict[str, Any]] = {} + for channel_section in channels: + channel = channel_section.get("channel") + if not isinstance(channel, int): + continue + summary = _waveform_summary(channel_section) + snapshot = _scope_status_data(channel_section) + frequency_hz, source, confidence, withheld_reason = _frequency_for_advice( + summary, + expected_frequencies.get(channel), + ) + vertical_scale = _recommended_vertical_scale( + summary, + snapshot, + target_vertical_divisions=target_vertical_divisions, + ) + time_range = ( + _recommended_time_range(frequency_hz, target_cycles=target_cycles) + if frequency_hz is not None + else None + ) + channel_profiles[channel] = { + "channel": channel, + "frequency_hz": frequency_hz, + "frequency_source": source, + "frequency_confidence": confidence, + "time_range_s": time_range, + "vertical_scale_v_per_div": vertical_scale, + } + if snapshot and snapshot.get("channel", {}).get("enabled") is False: + recommendations.append( + _command_recommendation( + "display_on", + "high", + channel, + "Channel display is off; enable it before human visual inspection.", + "display", + {"channel": channel, "state": "on"}, + ) + ) + if time_range is None and withheld_reason is not None: + # 低置信度测量又没有可用的期望频率时,明确说明为何不给时基建议 + recommendations.append( + { + "id": "timebase_advice_withheld", + "priority": "normal", + "channel": channel, + "action": "provide_expected_frequency", + "reason": withheld_reason, + "mutates_instrument_if_applied": False, + "raw_scpi": False, + } + ) + if time_range is not None or vertical_scale is not None: + reason = _focus_reason( + summary, + frequency_hz, + source, + confidence, + target_cycles=target_cycles, + ) + priority = "high" if _needs_focus(summary, channel, expected_frequencies) else "normal" + recommendations.append( + _command_recommendation( + "focus_channel", + priority, + channel, + reason, + "focus", + { + "channel": channel, + "time_range_s": time_range, + "vertical_scale_v_per_div": vertical_scale, + "frequency_confidence": confidence, + "hide_other_channels": False, + }, + ) + ) + span = _frequency_span(channel_profiles) + if span is not None and span["ratio_high_over_low"] > 10.0: + recommendations.append( + { + "id": "separate_timebase_profiles", + "priority": "high", + "action": "capture_or_observe_channels_separately", + "reason": ( + "Observed or expected channel frequencies span more than 10x; " + "do not judge every waveform shape on one timebase." + ), + "mutates_instrument_if_applied": False, + "raw_scpi": False, + "frequency_span": span, + "profiles": [ + profile + for _, profile in sorted(channel_profiles.items()) + if profile["time_range_s"] is not None + ], + } + ) + if not recommendations: + recommendations.append( + { + "id": "no_adjustment_needed", + "priority": "low", + "action": "keep_current_scope_settings", + "reason": "No obvious display or acquisition-window issue was found.", + "mutates_instrument_if_applied": False, + "raw_scpi": False, + } + ) + return recommendations + + +def _waveform_summary(channel_section: dict[str, Any]) -> dict[str, Any] | None: + waveform = channel_section.get("waveform", {}) + if waveform.get("status") != "ok": + return None + summary = waveform.get("data", {}).get("summary") + return summary if isinstance(summary, dict) else None + + +def _scope_status_data(channel_section: dict[str, Any]) -> dict[str, Any] | None: + status = channel_section.get("scope_status", {}) + data = status.get("data") + return data if isinstance(data, dict) else None + + +def _frequency_for_advice( + summary: dict[str, Any] | None, + expected_frequency_hz: float | None, +) -> tuple[float | None, str | None, str | None, str | None]: + """返回 (频率, 来源, 置信度, 撤回建议的原因)。 + + 低置信度的测量频率(``low_cycle_count`` 等质量告警)不能用来推导时基建议; + 此时优先回退到调用方提供的期望频率,没有期望频率就不给时基建议。 + """ + measured = _summary_frequency(summary) + if measured is not None and not _summary_frequency_low_confidence(summary): + return measured, "measured", "measured", None + if expected_frequency_hz is not None: + source = "expected" if measured is None else "expected_over_low_confidence_measurement" + return expected_frequency_hz, source, "configured", None + if measured is not None: + return ( + None, + None, + "low", + "measured frequency is low confidence (few cycles in window) and no expected frequency was provided", + ) + return None, None, None, None + + +def _summary_frequency(summary: dict[str, Any] | None) -> float | None: + if summary is None: + return None + value = summary.get("frequency_estimate_hz") + if not isinstance(value, (int, float)) or isinstance(value, bool) or value <= 0: + return None + return float(value) + + +def _summary_frequency_low_confidence(summary: dict[str, Any] | None) -> bool: + if summary is None: + return False + return any( + str(item).startswith("low_cycle_count") + for item in summary.get("quality_warnings", []) or [] + ) + + +def _recommended_time_range(frequency_hz: float, *, target_cycles: float) -> float: + return float(target_cycles / frequency_hz) + + +def _recommended_vertical_scale( + summary: dict[str, Any] | None, + snapshot: dict[str, Any] | None, + *, + target_vertical_divisions: float, +) -> float | None: + vpp = None if summary is None else summary.get("voltage_vpp_v") + if isinstance(vpp, (int, float)) and not isinstance(vpp, bool) and vpp > 0: + return float(vpp) / target_vertical_divisions + scale = None + if snapshot is not None: + scale = snapshot.get("channel", {}).get("scale_v_per_div") + if isinstance(scale, (int, float)) and not isinstance(scale, bool) and scale > 0: + return float(scale) + return None + + +def _needs_focus( + summary: dict[str, Any] | None, + channel: int, + expected_frequencies: dict[int, float], +) -> bool: + if channel in expected_frequencies and summary is None: + return True + if summary is None: + return False + cycles = summary.get("estimated_cycles") + if isinstance(cycles, (int, float)) and (cycles < 5.0 or cycles > 25.0): + return True + points_per_cycle = summary.get("points_per_cycle") + if isinstance(points_per_cycle, (int, float)) and points_per_cycle < 20.0: + return True + return bool(summary.get("quality_warnings")) + + +def _focus_reason( + summary: dict[str, Any] | None, + frequency_hz: float | None, + frequency_source: str | None, + frequency_confidence: str | None, + *, + target_cycles: float, +) -> str: + parts: list[str] = [] + if frequency_hz is not None: + parts.append( + f"use {frequency_source} frequency {frequency_hz:.6g} Hz " + f"(confidence={frequency_confidence}) to show about {target_cycles:.3g} cycles" + ) + if summary is not None: + cycles = summary.get("estimated_cycles") + if isinstance(cycles, (int, float)): + parts.append(f"current window contains about {cycles:.3g} cycles") + points = summary.get("points_per_cycle") + if isinstance(points, (int, float)): + parts.append(f"current sampling density is about {points:.3g} points/cycle") + return "; ".join(parts) if parts else "focus the selected channel for visual inspection" + + +def _frequency_span(profiles: dict[int, dict[str, Any]]) -> dict[str, Any] | None: + values = [ + (channel, profile["frequency_hz"]) + for channel, profile in profiles.items() + if isinstance(profile.get("frequency_hz"), (int, float)) and profile["frequency_hz"] > 0 + ] + if len(values) < 2: + return None + low_channel, low = min(values, key=lambda item: item[1]) + high_channel, high = max(values, key=lambda item: item[1]) + return { + "low_channel": low_channel, + "low_hz": low, + "high_channel": high_channel, + "high_hz": high, + "ratio_high_over_low": float(high / low), + } + + +def _command_recommendation( + recommendation_id: str, + priority: str, + channel: int, + reason: str, + command: str, + parameters: dict[str, Any], +) -> dict[str, Any]: + return { + "id": recommendation_id, + "priority": priority, + "channel": channel, + "action": f"scope.{command}", + "reason": reason, + "command": _command_text(command, parameters), + "parameters": parameters, + "mutates_instrument_if_applied": True, + "raw_scpi": False, + } + + +def _command_text(command: str, parameters: dict[str, Any]) -> str: + if command == "display": + return ( + "wavebench scope display " + f"--channel {parameters['channel']} {parameters['state']}" + ) + pieces = ["wavebench", "scope", "focus", "--channel", str(parameters["channel"])] + if parameters.get("time_range_s") is not None: + pieces.extend(["--time-range", f"{parameters['time_range_s']:.12g}"]) + if parameters.get("vertical_scale_v_per_div") is not None: + pieces.extend([ + "--vertical-scale", + f"{parameters['channel']}={parameters['vertical_scale_v_per_div']:.12g}", + ]) + if parameters.get("hide_other_channels"): + pieces.append("--hide-others") + return " ".join(pieces) + + +def _agent_hints( + observation: dict[str, Any], + recommendations: list[dict[str, Any]], +) -> list[str]: + hints = list(observation.get("agent_hints", [])) + if any(item["id"] == "separate_timebase_profiles" for item in recommendations): + hints.append("advise: run focus/observe per channel when frequencies differ greatly") + if any(item["id"] == "timebase_advice_withheld" for item in recommendations): + hints.append( + "advise: timebase advice withheld for at least one channel because the measured frequency " + "is low confidence and no expected frequency was provided" + ) + if observation.get("mutates_instrument"): + hints.append("advise: recommendations were computed from an explicit waveform read and were not applied") + else: + hints.append("advise: recommendations were computed without reading waveforms or changing instrument state") + return hints diff --git a/src/wavebench/services/agent_observe.py b/src/wavebench/services/agent_observe.py new file mode 100644 index 00000000..a5a1b6d4 --- /dev/null +++ b/src/wavebench/services/agent_observe.py @@ -0,0 +1,388 @@ +from __future__ import annotations + +from dataclasses import asdict +from pathlib import Path +from typing import Any + +from wavebench.config import load_config +from wavebench.data.expectations import ( + evaluate_waveform_expectation, + expectation_summary, + validate_expectation, +) +from wavebench.data.relationships import analyze_waveform_relationships +from wavebench.errors import ConfigError, WaveBenchError +from wavebench.instruments.models import WaveformData +from wavebench.logging import CommandLogger +from wavebench.services.scope_service import ScopeService + +# 读取波形可能造成的仪器状态影响。读取前不恢复原采集状态,调用方必须先确认。 +_WAVEFORM_STATE_EFFECTS = [ + "a running acquisition may be stopped", + "waveform transfer source/mode/format/points may be changed", + "some drivers may enable the requested analog channel display before fetching", + "the previous acquisition run state is not restored", +] + + +def scope_observe_payload( + *, + config_path: str | Path, + channel: int | None = None, + channels: tuple[int, ...] | None = None, + allow_50ohm: bool = False, + resource: str | None = None, +) -> dict[str, Any]: + """严格只读的示波器观察:IDN、每通道快照和高阻安全判断。 + + 该函数不读取波形,也不发送任何会改变仪器状态的命令,因此可以安全地通过 MCP 暴露。 + 需要波形、期望值检查或多通道关系时,请使用 ``scope_waveform_report_payload`` + (对应显式 CLI 命令 ``wavebench scope observe --fetch-waveform``)。 + """ + return _build_observation( + config_path=config_path, + channel=channel, + channels=channels, + allow_50ohm=allow_50ohm, + resource=resource, + fetch_waveform=False, + expectations=None, + ) + + +def scope_waveform_report_payload( + *, + config_path: str | Path, + channel: int | None = None, + channels: tuple[int, ...] | None = None, + allow_50ohm: bool = False, + expectations: dict[int, dict[str, Any]] | None = None, + resource: str | None = None, +) -> dict[str, Any]: + """显式读取波形并给出摘要、期望值检查和多通道关系。 + + 读取波形属于写操作:可能停止正在运行的采集、修改波形传输参数并打开通道显示。 + 所有输入(通道、期望值)都在任何仪器 I/O 之前完成校验,非法输入不会产生任何仪器写入。 + """ + normalized_expectations = _normalize_expectations(expectations) + return _build_observation( + config_path=config_path, + channel=channel, + channels=channels, + allow_50ohm=allow_50ohm, + resource=resource, + fetch_waveform=True, + expectations=normalized_expectations, + ) + + +def _build_observation( + *, + config_path: str | Path, + channel: int | None, + channels: tuple[int, ...] | None, + allow_50ohm: bool, + resource: str | None, + fetch_waveform: bool, + expectations: dict[int, dict[str, Any]] | None, +) -> dict[str, Any]: + config = load_config(config_path) + if resource: + config = config.with_resource(resource) + observed_channels = _scope_channels( + channel=channel, + channels=channels, + default_channel=config.scope.default_channel, + ) + unknown_expectation_channels = sorted( + item for item in (expectations or {}) if item not in observed_channels + ) + if unknown_expectation_channels: + raise ConfigError( + "scope observe expectation channels must be observed channels: " + f"{', '.join(str(item) for item in unknown_expectation_channels)}" + ) + service = ScopeService(config=config, logger=CommandLogger()) + sections: dict[str, Any] = {} + warnings: list[str] = [] + fetched_waveforms: dict[int, WaveformData] = {} + expectation_results: dict[int, dict[str, Any]] = {} + + sections["identity"] = _attempt(lambda: {"idn": service.idn()}, warnings=warnings, name="identity") + channel_sections = [ + _observe_channel( + service, + observed_channel, + fetch_waveform=fetch_waveform, + allow_50ohm=allow_50ohm, + warnings=warnings, + fetched_waveforms=fetched_waveforms, + expectations=expectations or {}, + expectation_results=expectation_results, + ) + for observed_channel in observed_channels + ] + first_channel = channel_sections[0] + sections["scope_status"] = first_channel["scope_status"] + sections["coupling"] = first_channel["coupling"] + + payload: dict[str, Any] = { + "status": "ok" if not warnings else "partial", + "read_only": not fetch_waveform, + "query_only": not fetch_waveform, + "mutates_instrument": fetch_waveform, + "raw_scpi": False, + "instrument_state_effects": list(_WAVEFORM_STATE_EFFECTS) if fetch_waveform else [], + "config": { + "path": str(config.source_path), + "scope_driver": config.scope.driver, + "resource": config.connection.resource, + "backend": config.connection.backend, + "default_channel": config.scope.default_channel, + "waveform_points": config.waveform.points, + }, + "observation": { + "instrument": "scope", + "channel": observed_channels[0], + "channels": list(observed_channels), + "fetch_waveform": fetch_waveform, + "allow_50ohm": allow_50ohm, + }, + **sections, + "channels": channel_sections, + "warnings": warnings, + } + if fetch_waveform: + # 每个通道各自打开 session,波形不保证来自同一次 acquisition;跨采集的时序关系不成立。 + payload["waveform_source"] = { + "same_acquisition": False, + "reason": "channels are fetched channel-by-channel, not in one acquisition", + } + payload["relationships"] = ( + analyze_waveform_relationships(fetched_waveforms, same_acquisition=False) + if len(fetched_waveforms) >= 2 + else [] + ) + payload["expectations"] = expectation_summary(expectation_results) + payload["agent_hints"] = _agent_hints( + sections, + warnings, + channel_sections=channel_sections, + fetched_waveforms=fetched_waveforms, + expectation_results=expectation_results, + fetch_waveform=fetch_waveform, + ) + return payload + + +def _scope_channels( + *, + channel: int | None, + channels: tuple[int, ...] | None, + default_channel: int, +) -> tuple[int, ...]: + if channel is not None and channels is not None: + raise ConfigError("scope observe accepts either channel or channels, not both") + candidates = channels if channels is not None else (default_channel if channel is None else channel,) + if not candidates: + raise ConfigError("scope observe channels must not be empty") + for candidate in candidates: + if isinstance(candidate, bool) or not isinstance(candidate, int) or candidate < 1: + raise ConfigError("scope observe channel must be a positive integer") + if len(set(candidates)) != len(candidates): + raise ConfigError("scope observe channels must be unique") + return candidates + + +def _normalize_expectations( + expectations: dict[int, dict[str, Any]] | None, +) -> dict[int, dict[str, Any]]: + if expectations is None: + return {} + normalized: dict[int, dict[str, Any]] = {} + for channel, expectation in expectations.items(): + if isinstance(channel, bool) or not isinstance(channel, int) or channel < 1: + raise ConfigError("scope observe expectation channel must be a positive integer") + normalized[channel] = validate_expectation(expectation) + return normalized + + +def _observe_channel( + service: ScopeService, + channel: int, + *, + fetch_waveform: bool, + allow_50ohm: bool, + warnings: list[str], + fetched_waveforms: dict[int, WaveformData], + expectations: dict[int, dict[str, Any]], + expectation_results: dict[int, dict[str, Any]], +) -> dict[str, Any]: + section: dict[str, Any] = { + "channel": channel, + "scope_status": _attempt( + lambda: asdict(service.status(channel=channel)), + warnings=warnings, + name=f"ch{channel}_scope_status", + ), + "coupling": _attempt( + lambda: _coupling_payload(service, channel, allow_50ohm=allow_50ohm), + warnings=warnings, + name=f"ch{channel}_coupling", + ), + } + if not fetch_waveform: + return section + section["waveform"] = _attempt( + lambda: _waveform_payload( + service, + channel, + allow_50ohm=allow_50ohm, + fetched_waveforms=fetched_waveforms, + ), + warnings=warnings, + name=f"ch{channel}_waveform", + ) + if channel in expectations and section["waveform"]["status"] == "ok": + result = _attempt( + lambda: evaluate_waveform_expectation( + fetched_waveforms[channel], + expectations[channel], + ), + warnings=warnings, + name=f"ch{channel}_expectation", + ) + section["expectation"] = result + elif channel in expectations: + section["expectation"] = { + "status": "unavailable", + "reason": "waveform unavailable", + } + warnings.append(f"ch{channel}_expectation_unavailable: waveform unavailable") + if channel in expectations: + result = section["expectation"] + expectation_results[channel] = ( + result["data"] if result["status"] == "ok" + else {**result, "channel": channel, "checks": []} + ) + return section + + +def _attempt(call, *, warnings: list[str], name: str) -> dict[str, Any]: + try: + return {"status": "ok", "data": call()} + except WaveBenchError as exc: + warnings.append(f"{name}_unavailable: {exc}") + return { + "status": "unavailable", + "error": {"type": type(exc).__name__, "message": str(exc)}, + } + except Exception as exc: + warnings.append(f"{name}_unavailable: {type(exc).__name__}: {exc}") + return { + "status": "unavailable", + "error": {"type": type(exc).__name__, "message": str(exc)}, + } + + +def _coupling_payload( + service: ScopeService, + channel: int, + *, + allow_50ohm: bool, +) -> dict[str, Any]: + coupling = service.require_high_impedance(channel, allow_50ohm=allow_50ohm) + return { + "channel": channel, + "coupling": coupling, + "accepted_for_capture": True, + } + + +def _waveform_payload( + service: ScopeService, + channel: int, + *, + allow_50ohm: bool, + fetched_waveforms: dict[int, WaveformData], +) -> dict[str, Any]: + service.require_high_impedance(channel, allow_50ohm=allow_50ohm) + waveform = service.fetch_waveform(channel=channel) + fetched_waveforms[channel] = waveform + return { + "channel": channel, + "summary": waveform.summary( + expected_frequency_hz=service.config.waveform.expected_frequency_hz, + frequency_tolerance_ratio=service.config.waveform.frequency_tolerance_ratio, + ), + "raw_samples_included": False, + } + + +def _agent_hints( + sections: dict[str, Any], + warnings: list[str], + *, + channel_sections: list[dict[str, Any]], + fetched_waveforms: dict[int, WaveformData], + expectation_results: dict[int, dict[str, Any]], + fetch_waveform: bool, +) -> list[str]: + hints: list[str] = [] + if not fetch_waveform: + hints.append( + "read-only observation: waveforms were not read; use `wavebench scope observe --fetch-waveform` " + "when waveform summaries, expectations or multi-channel relationships are required" + ) + for channel_section in channel_sections: + waveform = channel_section.get("waveform", {}) + if waveform.get("status") != "ok": + continue + channel = channel_section.get("channel") + summary = waveform.get("data", {}).get("summary", {}) + for warning in summary.get("quality_warnings", []) or []: + hints.append(f"CH{channel}_waveform_quality_warning: {warning}") + cycles = summary.get("estimated_cycles") + if isinstance(cycles, (int, float)) and cycles < 5: + hints.append(f"CH{channel}: consider capturing a wider time window for robust periodic analysis") + if len(fetched_waveforms) >= 2: + hints.append( + "waveforms were read channel-by-channel, so they are not from one acquisition; " + "timing relationships (phase/correlation/intersections) were skipped. " + "Use `wavebench scope capture --channel ... --synchronized` for driver-proven single-acquisition capture" + ) + frequencies = _trusted_frequencies(fetched_waveforms) + if len(frequencies) >= 2 and min(frequencies) > 0 and max(frequencies) / min(frequencies) > 10: + hints.append( + "multi_channel_frequency_span_large: use separate time windows/profiles before judging waveform shape across channels" + ) + for channel, result in sorted(expectation_results.items()): + if result["status"] in {"warn", "fail"}: + hints.append(f"CH{channel}_expectation_{result['status']}: inspect expectation checks") + elif result["status"] == "skipped": + hints.append(f"CH{channel}_expectation_skipped: expectation contains no checkable metric") + elif result["status"] == "unavailable": + hints.append(f"CH{channel}_expectation_unavailable: acceptance could not be evaluated") + if sections.get("scope_status", {}).get("status") == "unavailable": + hints.append("driver lacks scope.snapshot or the status query failed; use identity cautiously") + if sections.get("coupling", {}).get("status") == "unavailable": + hints.append("do not run capture until input coupling safety is confirmed") + if warnings: + hints.append("treat this observation as partial and avoid state-changing actions") + return hints + + +def _trusted_frequencies(waveforms: dict[int, WaveformData]) -> list[float]: + frequencies: list[float] = [] + for waveform in waveforms.values(): + summary = waveform.summary() + value = summary.get("frequency_estimate_hz") + if not isinstance(value, (int, float)) or value <= 0: + continue + if any( + str(item).startswith("low_cycle_count") + for item in summary.get("quality_warnings", []) or [] + ): + continue + frequencies.append(float(value)) + return frequencies diff --git a/tests/test_agent_advise.py b/tests/test_agent_advise.py new file mode 100644 index 00000000..5d5df8c8 --- /dev/null +++ b/tests/test_agent_advise.py @@ -0,0 +1,250 @@ +import shlex +from pathlib import Path +from tempfile import TemporaryDirectory +from unittest.mock import patch + +import pytest + +from wavebench.cli import _scope_focus_request +from wavebench.cli_parser import build_parser +from wavebench.errors import ConfigError +from wavebench.services.agent_advise import ( + _command_text, + scope_advise_from_observation, + scope_advise_payload, +) + + +def _write_config(root: Path) -> Path: + path = root / "wavebench.toml" + path.write_text( + """ +[connection] +resource = "TCPIP::scope::INSTR" + +[scope] +driver = "ds1104" +default_channel = 1 +""", + encoding="utf-8", + ) + return path + + +class _FakeScopeService: + def __init__(self, *, config, logger): + self.config = config + + def idn(self): + return "RIGOL TECHNOLOGIES,DS1104Z Plus,123,1.0" + + def status(self, channel): + return _snapshot(channel) + + def require_high_impedance(self, channel, *, allow_50ohm=False): + return "DC" + + +def _snapshot(channel: int): + from wavebench.instruments.models import ( + ScopeAnalogChannelSnapshot, + ScopeEdgeTriggerSnapshot, + ScopeHealthSnapshot, + ScopeIdentitySnapshot, + ScopeProbeSnapshot, + ScopeSnapshot, + ScopeTimebaseSnapshot, + ScopeWaveformMetadataSnapshot, + ) + + return ScopeSnapshot( + identity=ScopeIdentitySnapshot("RIGOL", "DS1104Z", "123", "1.0", ()), + health=ScopeHealthSnapshot(0, 0, 0, 1, 1, 1_000_000.0, False, False), + channel=ScopeAnalogChannelSnapshot( + channel, True, "DC", 8.0, 1.0, 0.0, 0.0, None, "NORM", 0.0, "", False, False, "SAMPLE" + ), + timebase=ScopeTimebaseSnapshot(0.001, 12, 0.0, 0.0012, 50.0, 0.0001, False), + probe=ScopeProbeSnapshot(channel, 10.0, None, None, 1_000_000.0, "P10", "PASSIVE"), + waveform=ScopeWaveformMetadataSnapshot( + channel, -0.0005, 0.0005, 1000, 1, 1e-6, -0.0005, 0.001, 0.0, 8 + ), + trigger=ScopeEdgeTriggerSnapshot("EDGE", channel, "AUTO", "POS", "DC", 0.0, "AUTO", "OFF", 1e-6), + ) + + +def _observation(*, fetch_waveform: bool, measured_frequency: dict[int, float] | None = None) -> dict: + def channel_section(channel: int, frequency: float | None, warnings: list[str]) -> dict: + section = { + "channel": channel, + "scope_status": { + "status": "ok", + "data": {"channel": {"enabled": True, "scale_v_per_div": 1.0}}, + }, + "coupling": {"status": "ok", "data": {"channel": channel, "coupling": "DC", "accepted_for_capture": True}}, + } + if frequency is not None: + section["waveform"] = { + "status": "ok", + "data": { + "summary": { + "frequency_estimate_hz": frequency, + "estimated_cycles": 2.4 if warnings else 120.0, + "points_per_cycle": 500.0, + "voltage_vpp_v": 1.0, + "quality_warnings": warnings, + } + }, + } + return section + + measured = measured_frequency or {} + channels = [ + channel_section( + 1, + measured.get(1), + ["low_cycle_count: 2.4"] if 1 in measured and len(measured) == 1 else [], + ), + channel_section(2, measured.get(2), []), + ] + return { + "status": "ok", + "read_only": not fetch_waveform, + "query_only": not fetch_waveform, + "mutates_instrument": fetch_waveform, + "raw_scpi": False, + "instrument_state_effects": ["a running acquisition may be stopped"] if fetch_waveform else [], + "observation": {"channel": 1, "channels": [1, 2], "fetch_waveform": fetch_waveform}, + "channels": channels, + "relationships": [], + "warnings": [], + "agent_hints": [], + } + + +def test_scope_advise_payload_is_read_only_and_uses_expected_frequencies(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + payload = scope_advise_payload( + config_path=config, + channels=(1, 2), + expected_frequencies_hz={1: 1000.0, 2: 50000.0}, + ) + + assert payload["read_only"] is True + assert payload["query_only"] is True + assert payload["mutates_instrument"] is False + assert payload["applies_recommendations"] is False + assert payload["instrument_state_effects"] == [] + focus = [item for item in payload["recommendations"] if item["id"] == "focus_channel"] + assert [item["channel"] for item in focus] == [1, 2] + assert focus[0]["parameters"]["time_range_s"] == pytest.approx(0.01) + assert focus[0]["parameters"]["frequency_confidence"] == "configured" + assert focus[1]["parameters"]["time_range_s"] == pytest.approx(0.0002) + span = payload["recommendations"][-1] + assert span["id"] == "separate_timebase_profiles" + assert span["frequency_span"]["ratio_high_over_low"] == pytest.approx(50.0) + + +def test_scope_advise_prefers_expected_frequency_over_low_confidence_measurement(): + observation = _observation(fetch_waveform=True, measured_frequency={1: 1000.0}) + + payload = scope_advise_from_observation( + observation, + expected_frequencies_hz={1: 2000.0}, + ) + + focus = [item for item in payload["recommendations"] if item["id"] == "focus_channel"] + assert focus[0]["parameters"]["time_range_s"] == pytest.approx(10.0 / 2000.0) + assert ( + focus[0]["parameters"]["frequency_confidence"] == "configured" + ) + assert focus[0]["parameters"]["time_range_s"] is not None + + +def test_scope_advise_withholds_timebase_advice_when_measurement_is_low_confidence(): + observation = _observation(fetch_waveform=True, measured_frequency={1: 1000.0}) + + payload = scope_advise_from_observation(observation) + + # 低置信度测量频率没有期望频率可回退时,不给基于该频率的时基建议 + withheld = [item for item in payload["recommendations"] if item["id"] == "timebase_advice_withheld"] + assert len(withheld) == 1 + assert withheld[0]["channel"] == 1 + assert "low confidence" in withheld[0]["reason"] + # 只允许保留与频率无关的垂直档位建议,时基建议必须为空 + focus = [ + item + for item in payload["recommendations"] + if item["id"] == "focus_channel" and item["channel"] == 1 + ] + assert [item["parameters"]["time_range_s"] for item in focus] == [None] + assert any("timebase advice withheld" in hint for hint in payload["agent_hints"]) + + +def test_scope_advise_uses_trusted_measured_frequency_when_confidence_is_good(): + observation = _observation(fetch_waveform=True, measured_frequency={1: 1000.0, 2: 50000.0}) + + payload = scope_advise_from_observation(observation) + + focus = [item for item in payload["recommendations"] if item["id"] == "focus_channel"] + assert focus[0]["parameters"]["frequency_confidence"] == "measured" + assert focus[0]["parameters"]["time_range_s"] == pytest.approx(0.01) + + +@pytest.mark.parametrize("value", [0.0, -1.0, float("nan"), float("inf")]) +def test_scope_advise_rejects_invalid_targets(value): + with pytest.raises(ConfigError, match="target_cycles"): + scope_advise_payload(config_path="wavebench.toml", target_cycles=value) + + +@pytest.mark.parametrize("value", [0.0, -1.0, float("nan"), float("inf")]) +def test_scope_advise_rejects_invalid_expected_frequencies(value): + with pytest.raises(ConfigError, match="expected frequency"): + scope_advise_from_observation(_observation(fetch_waveform=False), expected_frequencies_hz={1: value}) + + +@pytest.mark.parametrize("fetch_waveform", [False, True]) +def test_advice_commands_parse_with_real_cli(fetch_waveform): + observation = _observation( + fetch_waveform=fetch_waveform, + measured_frequency={1: 1000.0, 2: 50000.0} if fetch_waveform else None, + ) + observation["channels"][0]["scope_status"]["data"]["channel"]["enabled"] = False + payload = scope_advise_from_observation( + observation, expected_frequencies_hz={1: 1000.0, 2: 50000.0}, + ) + parser = build_parser() + commands = [item for item in payload["recommendations"] if "command" in item] + assert len(commands) == 3 + for recommendation in commands: + args = parser.parse_args(shlex.split(recommendation["command"])[1:]) + parameters = recommendation["parameters"] + if recommendation["action"] == "scope.display": + assert args.channel == parameters["channel"] + assert args.state == "on" + else: + request = _scope_focus_request(args) + assert request.channels == (parameters["channel"],) + assert request.time_range_s == pytest.approx(parameters["time_range_s"]) + assert len(request.vertical_scales) == 1 + assert request.vertical_scales[0].channel == parameters["channel"] + assert request.vertical_scales[0].scale_v_per_div == pytest.approx( + parameters["vertical_scale_v_per_div"] + ) + assert request.hide_others is False + + +@pytest.mark.parametrize("hide_others", [False, True]) +def test_focus_command_hide_others_uses_real_cli_flag(hide_others): + command = _command_text("focus", { + "channel": 2, + "vertical_scale_v_per_div": 0.25, + "hide_other_channels": hide_others, + }) + args = build_parser().parse_args(shlex.split(command)[1:]) + request = _scope_focus_request(args) + + assert request.channels == (2,) + assert request.vertical_scales[0].scale_v_per_div == 0.25 + assert request.hide_others is hide_others diff --git a/tests/test_agent_observe.py b/tests/test_agent_observe.py new file mode 100644 index 00000000..38c351c9 --- /dev/null +++ b/tests/test_agent_observe.py @@ -0,0 +1,296 @@ +from pathlib import Path +from tempfile import TemporaryDirectory +from unittest.mock import patch + +import numpy as np +import pytest + +from wavebench.errors import ConfigError +from wavebench.data.expectations import evaluate_waveform_expectation +from wavebench.instruments.models import ( + ScopeAnalogChannelSnapshot, + ScopeEdgeTriggerSnapshot, + ScopeHealthSnapshot, + ScopeIdentitySnapshot, + ScopeProbeSnapshot, + ScopeSnapshot, + ScopeTimebaseSnapshot, + ScopeWaveformMetadataSnapshot, + WaveformData, + WaveformHeader, +) +from wavebench.services.agent_observe import ( + scope_observe_payload, + scope_waveform_report_payload, +) + + +def _write_config(root: Path) -> Path: + path = root / "wavebench.toml" + path.write_text( + """ +[connection] +resource = "TCPIP::scope::INSTR" + +[scope] +driver = "ds1104" +default_channel = 1 + +[waveform] +points = "def" +""", + encoding="utf-8", + ) + return path + + +def _snapshot(channel: int) -> ScopeSnapshot: + return ScopeSnapshot( + identity=ScopeIdentitySnapshot("RIGOL", "DS1104Z", "123", "1.0", ()), + health=ScopeHealthSnapshot(0, 0, 0, 1, 1, 1_000_000.0, False, False), + channel=ScopeAnalogChannelSnapshot( + channel, + True, + "DC", + 8.0, + 1.0, + 0.0, + 0.0, + None, + "NORM", + 0.0, + "", + False, + False, + "SAMPLE", + ), + timebase=ScopeTimebaseSnapshot(0.001, 12, 0.0, 0.0012, 50.0, 0.0001, False), + probe=ScopeProbeSnapshot(channel, 10.0, None, None, 1_000_000.0, "P10", "PASSIVE"), + waveform=ScopeWaveformMetadataSnapshot( + channel, + -0.0005, + 0.0005, + 1000, + 1, + 1e-6, + -0.0005, + 0.001, + 0.0, + 8, + ), + trigger=ScopeEdgeTriggerSnapshot("EDGE", channel, "AUTO", "POS", "DC", 0.0, "AUTO", "OFF", 1e-6), + ) + + +class _FakeScopeService: + """记录是否真的读取过波形,用于断言只读路径不碰仪器写路径。""" + + instances: list["_FakeScopeService"] = [] + + def __init__(self, *, config, logger): + self.config = config + self.fetched_channels: list[int] = [] + self.allow_50ohm_seen: list[bool] = [] + _FakeScopeService.instances.append(self) + + def idn(self): + return "RIGOL TECHNOLOGIES,DS1104Z Plus,123,1.0" + + def status(self, channel): + return _snapshot(channel) + + def require_high_impedance(self, channel, *, allow_50ohm=False): + self.allow_50ohm_seen.append(allow_50ohm) + return "DC" + + def fetch_waveform(self, channel): + self.fetched_channels.append(channel) + times = np.linspace(0.0, 0.005, 2000) + return WaveformData( + channel=channel, + header=WaveformHeader(x_start=0.0, x_stop=0.005, points=2000), + voltages_v=np.sin(2 * np.pi * 1000 * times), + ) + + +@pytest.fixture(autouse=True) +def _reset_instances(): + _FakeScopeService.instances = [] + yield + _FakeScopeService.instances = [] + + +def test_scope_observe_payload_is_strictly_read_only(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + payload = scope_observe_payload(config_path=config, channel=2) + + assert payload["status"] == "ok" + assert payload["read_only"] is True + assert payload["query_only"] is True + assert payload["mutates_instrument"] is False + assert payload["raw_scpi"] is False + assert payload["instrument_state_effects"] == [] + assert payload["observation"]["channel"] == 2 + assert payload["observation"]["channels"] == [2] + assert payload["identity"]["data"]["idn"].startswith("RIGOL") + assert payload["scope_status"]["data"]["channel"]["channel"] == 2 + assert payload["coupling"]["data"]["accepted_for_capture"] is True + # 只读路径不得读取波形,也不得暴露波形/关系字段 + assert "waveform" not in payload["channels"][0] + assert "relationships" not in payload + assert "expectations" not in payload + assert _FakeScopeService.instances[0].fetched_channels == [] + assert any("read-only observation" in hint for hint in payload["agent_hints"]) + + +def test_scope_observe_payload_supports_multiple_channels(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + payload = scope_observe_payload(config_path=config, channels=(1, 2)) + + assert payload["observation"]["channel"] == 1 + assert payload["observation"]["channels"] == [1, 2] + assert [item["channel"] for item in payload["channels"]] == [1, 2] + + +def test_scope_observe_payload_passes_allow_50ohm_to_the_safety_check(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + scope_observe_payload(config_path=config, channel=1, allow_50ohm=True) + + assert _FakeScopeService.instances[0].allow_50ohm_seen == [True] + + +def test_scope_waveform_report_is_an_explicit_write_path(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + payload = scope_waveform_report_payload(config_path=config, channel=1) + + assert payload["read_only"] is False + assert payload["query_only"] is False + assert payload["mutates_instrument"] is True + assert payload["instrument_state_effects"] + assert any("acquisition may be stopped" in item for item in payload["instrument_state_effects"]) + assert payload["waveform_source"]["same_acquisition"] is False + assert payload["channels"][0]["waveform"]["data"]["summary"]["samples"] == 2000 + assert _FakeScopeService.instances[0].fetched_channels == [1] + assert payload["expectations"] == {"status": "skipped", "channels": {}} + + +def test_scope_waveform_report_marks_multi_channel_timing_analysis_as_skipped(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + payload = scope_waveform_report_payload(config_path=config, channels=(1, 2)) + + assert payload["waveform_source"]["same_acquisition"] is False + relationship = payload["relationships"][0] + assert relationship["channels"] == [1, 2] + assert relationship["correlation"] == {"status": "skipped", "reason": "not_same_acquisition"} + assert relationship["phase_degrees_at_left_frequency"] is None + assert any("not from one acquisition" in hint for hint in payload["agent_hints"]) + + +def test_scope_waveform_report_evaluates_channel_expectations(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + payload = scope_waveform_report_payload( + config_path=config, + channel=1, + expectations={1: {"frequency_hz": 1000.0, "frequency_tolerance_ratio": 0.05}}, + ) + + assert payload["expectations"]["status"] == "pass" + assert payload["channels"][0]["expectation"]["data"]["checks"][0]["metric"] == "frequency_hz" + + +@pytest.mark.parametrize("failed_channels", [(2,), (1, 2)]) +@pytest.mark.parametrize("failure", ["waveform", "coupling", "evaluation"]) +def test_scope_waveform_report_preserves_unavailable_expectations(failed_channels, failure): + class FailingScopeService(_FakeScopeService): + def fetch_waveform(self, channel): + if failure == "waveform" and channel in failed_channels: + raise ConfigError("waveform read failed") + return super().fetch_waveform(channel) + + def require_high_impedance(self, channel, *, allow_50ohm=False): + if failure == "coupling" and channel in failed_channels: + raise ConfigError("coupling safety unavailable") + return super().require_high_impedance(channel, allow_50ohm=allow_50ohm) + + def evaluate(waveform, expectation): + if failure == "evaluation" and waveform.channel in failed_channels: + raise ConfigError("expectation evaluation failed") + return evaluate_waveform_expectation(waveform, expectation) + + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with ( + patch("wavebench.services.agent_observe.ScopeService", FailingScopeService), + patch("wavebench.services.agent_observe.evaluate_waveform_expectation", evaluate), + ): + payload = scope_waveform_report_payload( + config_path=config, + channels=(1, 2), + expectations={1: {"vpp_v": 2.0}, 2: {"vpp_v": 2.0}}, + ) + + expected_statuses = { + str(channel): "unavailable" if channel in failed_channels else "pass" + for channel in (1, 2) + } + assert payload["status"] == "partial" + assert payload["expectations"] == { + "status": "unavailable" if len(failed_channels) == 2 else "partial", + "channels": expected_statuses, + } + for channel in failed_channels: + assert payload["channels"][channel - 1]["expectation"]["status"] == "unavailable" + assert any(f"CH{channel}_expectation_unavailable" in hint for hint in payload["agent_hints"]) + + +def test_scope_waveform_report_rejects_invalid_expectation_before_any_instrument_io(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + with pytest.raises(ConfigError, match="unknown expectation field"): + scope_waveform_report_payload( + config_path=config, + channel=1, + expectations={1: {"frequncy_hz": 1000}}, + ) + + # 校验失败时连 service 都不应创建,更不会产生仪器写入 + assert _FakeScopeService.instances == [] + + +def test_scope_waveform_report_rejects_expectation_for_unobserved_channel(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + with pytest.raises(ConfigError, match="expectation channels must be observed channels"): + scope_waveform_report_payload( + config_path=config, + channel=1, + expectations={2: {"frequency_hz": 1000.0}}, + ) + + +def test_scope_observe_rejects_invalid_channel(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with pytest.raises(ConfigError, match="positive integer"): + scope_observe_payload(config_path=config, channel=0) + + +def test_scope_observe_rejects_ambiguous_channel_arguments(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + with pytest.raises(ConfigError, match="either channel or channels"): + scope_observe_payload(config_path=config, channel=1, channels=(2,)) diff --git a/tests/test_mcp_http.py b/tests/test_mcp_http.py index 43538327..e507c1d3 100644 --- a/tests/test_mcp_http.py +++ b/tests/test_mcp_http.py @@ -5,7 +5,9 @@ import urllib.error import urllib.request from pathlib import Path +from types import SimpleNamespace from tempfile import TemporaryDirectory +from unittest.mock import patch from wavebench import __version__ from wavebench.cli import build_parser @@ -162,10 +164,26 @@ def test_tools_lists_only_read_only_mvp_tools(self): self.assertEqual(status, 200) names = {tool["name"] for tool in payload["tools"]} - self.assertEqual(names, {"run.schema", "run.check", "capture.inspect"}) + self.assertEqual( + names, + { + "run.schema", + "run.check", + "capture.inspect", + "scope.observe", + "doctor.config", + "scope.advise", + }, + ) self.assertFalse(any("raw" in name.lower() for name in names)) self.assertFalse(any("output" in name.lower() for name in names)) self.assertFalse(any(name.lower().endswith((".on", ".off")) for name in names)) + by_name = {tool["name"]: tool for tool in payload["tools"]} + # 所有 MCP 工具都必须是纯只读:不改变仪器状态 + for name, tool in by_name.items(): + self.assertTrue(tool["read_only"], name) + self.assertFalse(tool["mutates_instrument"], name) + self.assertEqual(tool["instrument_state_effects"], [], name) def test_call_run_schema_succeeds(self): with TemporaryDirectory() as tmp: @@ -335,7 +353,17 @@ def test_mcp_jsonrpc_tools_list_and_call(self): ) self.assertEqual(status, 200) names = {tool["name"] for tool in listed["result"]["tools"]} - self.assertEqual(names, {"run.schema", "run.check", "capture.inspect"}) + self.assertEqual( + names, + { + "run.schema", + "run.check", + "capture.inspect", + "scope.observe", + "doctor.config", + "scope.advise", + }, + ) status, called = self._request( server, @@ -353,6 +381,130 @@ def test_mcp_jsonrpc_tools_list_and_call(self): self.assertEqual(called["result"]["structuredContent"]["status"], "ok") self.assertEqual(called["result"]["content"][0]["type"], "text") + def test_call_scope_observe_succeeds_with_structured_read_only_payload(self): + with TemporaryDirectory() as tmp: + root = Path(tmp) + config = self._write_config(root) + server = self._start_server(config) + + with patch( + "wavebench.mcp_http.scope_observe_payload", + return_value={ + "status": "ok", + "read_only": True, + "mutates_instrument": False, + "raw_scpi": False, + "observation": {"channel": 2, "channels": [2, 3]}, + }, + ) as observe: + status, payload = self._request( + server, + "POST", + "/call", + token="test-token", + body={ + "tool": "scope.observe", + "arguments": {"channels": [2, 3]}, + }, + ) + + self.assertEqual(status, 200) + self.assertEqual(payload["result"]["status"], "ok") + self.assertTrue(payload["result"]["read_only"]) + self.assertFalse(payload["result"]["mutates_instrument"]) + observe.assert_called_once() + self.assertIsNone(observe.call_args.kwargs["channel"]) + self.assertEqual(observe.call_args.kwargs["channels"], (2, 3)) + # 只读工具不接受任何波形或期望值参数 + self.assertNotIn("fetch_waveform", observe.call_args.kwargs) + self.assertNotIn("expectations", observe.call_args.kwargs) + + def test_call_scope_observe_rejects_non_integer_channel(self): + with TemporaryDirectory() as tmp: + server = self._start_server(self._write_config(Path(tmp))) + + with self.assertRaises(urllib.error.HTTPError) as caught: + self._request( + server, + "POST", + "/call", + token="test-token", + body={ + "tool": "scope.observe", + "arguments": {"channel": "1"}, + }, + ) + + self.assertEqual(caught.exception.code, 400) + + def test_call_scope_advise_succeeds_without_applying_recommendations(self): + with TemporaryDirectory() as tmp: + root = Path(tmp) + config = self._write_config(root) + server = self._start_server(config) + + with patch( + "wavebench.mcp_http.scope_advise_payload", + return_value={ + "status": "ok", + "read_only": True, + "mutates_instrument": False, + "raw_scpi": False, + "applies_recommendations": False, + "recommendations": [{"id": "focus_channel"}], + }, + ) as advise: + status, payload = self._request( + server, + "POST", + "/call", + token="test-token", + body={ + "tool": "scope.advise", + "arguments": { + "channels": [1, 2], + "target_cycles": 8, + "expected_frequencies_hz": {"1": 1000}, + }, + }, + ) + + self.assertEqual(status, 200) + self.assertFalse(payload["result"]["applies_recommendations"]) + advise.assert_called_once() + self.assertEqual(advise.call_args.kwargs["channels"], (1, 2)) + self.assertEqual(advise.call_args.kwargs["target_cycles"], 8.0) + self.assertEqual(advise.call_args.kwargs["expected_frequencies_hz"], {1: 1000.0}) + + def test_call_doctor_config_returns_structured_records(self): + with TemporaryDirectory() as tmp: + root = Path(tmp) + config = self._write_config(root) + server = self._start_server(config) + record = SimpleNamespace( + severity="ok", + target="scope", + driver="ds1104", + resource="TCPIP::scope::INSTR", + idn="RIGOL,DS1104Z,123,1.0", + message="reachable", + suggestion="", + ) + + with patch("wavebench.mcp_http.doctor_records", return_value=[record]): + status, payload = self._request( + server, + "POST", + "/call", + token="test-token", + body={"tool": "doctor.config", "arguments": {"timeout_ms": 1000}}, + ) + + self.assertEqual(status, 200) + self.assertEqual(payload["result"]["status"], "ok") + self.assertEqual(payload["result"]["records"][0]["target"], "scope") + self.assertFalse(payload["result"]["mutates_instrument"]) + def test_mcp_jsonrpc_requires_token(self): with TemporaryDirectory() as tmp: server = self._start_server(self._write_config(Path(tmp))) diff --git a/tests/test_scope_observe_cli.py b/tests/test_scope_observe_cli.py new file mode 100644 index 00000000..83fd6a57 --- /dev/null +++ b/tests/test_scope_observe_cli.py @@ -0,0 +1,248 @@ +import io +import json +from contextlib import redirect_stderr, redirect_stdout +from pathlib import Path +from tempfile import TemporaryDirectory +from unittest.mock import patch + +import numpy as np +import pytest + +from wavebench.cli import _run_scope_observe, main +from wavebench.cli_parser import build_parser +from wavebench.errors import ConfigError +from wavebench.instruments.models import ( + ScopeAnalogChannelSnapshot, + ScopeEdgeTriggerSnapshot, + ScopeHealthSnapshot, + ScopeIdentitySnapshot, + ScopeProbeSnapshot, + ScopeSnapshot, + ScopeTimebaseSnapshot, + ScopeWaveformMetadataSnapshot, + WaveformData, + WaveformHeader, +) + + +def _write_config(root: Path) -> Path: + path = root / "wavebench.toml" + path.write_text( + """ +[connection] +resource = "TCPIP::scope::INSTR" + +[scope] +driver = "ds1104" +default_channel = 1 + +[waveform] +points = "def" +""", + encoding="utf-8", + ) + return path + + +def _snapshot(channel: int) -> ScopeSnapshot: + return ScopeSnapshot( + identity=ScopeIdentitySnapshot("RIGOL", "DS1104Z", "123", "1.0", ()), + health=ScopeHealthSnapshot(0, 0, 0, 1, 1, 1_000_000.0, False, False), + channel=ScopeAnalogChannelSnapshot( + channel, True, "DC", 8.0, 1.0, 0.0, 0.0, None, "NORM", 0.0, "", False, False, "SAMPLE" + ), + timebase=ScopeTimebaseSnapshot(0.001, 12, 0.0, 0.0012, 50.0, 0.0001, False), + probe=ScopeProbeSnapshot(channel, 10.0, None, None, 1_000_000.0, "P10", "PASSIVE"), + waveform=ScopeWaveformMetadataSnapshot( + channel, -0.0005, 0.0005, 1000, 1, 1e-6, -0.0005, 0.001, 0.0, 8 + ), + trigger=ScopeEdgeTriggerSnapshot("EDGE", channel, "AUTO", "POS", "DC", 0.0, "AUTO", "OFF", 1e-6), + ) + + +class _FakeScopeService: + instances: list["_FakeScopeService"] = [] + + def __init__(self, *, config, logger): + self.config = config + self.fetched_channels: list[int] = [] + _FakeScopeService.instances.append(self) + + def idn(self): + return "RIGOL TECHNOLOGIES,DS1104Z Plus,123,1.0" + + def status(self, channel): + return _snapshot(channel) + + def require_high_impedance(self, channel, *, allow_50ohm=False): + return "DC" + + def fetch_waveform(self, channel): + self.fetched_channels.append(channel) + times = np.linspace(0.0, 0.005, 2000) + return WaveformData( + channel=channel, + header=WaveformHeader(x_start=0.0, x_stop=0.005, points=2000), + voltages_v=np.sin(2 * np.pi * 1000 * times), + ) + + +def _run(argv: list[str]) -> tuple[int, str, str]: + _FakeScopeService.instances = [] + stdout = io.StringIO() + stderr = io.StringIO() + with patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService): + with redirect_stdout(stdout), redirect_stderr(stderr): + code = main(argv) + return code, stdout.getvalue(), stderr.getvalue() + + +def test_scope_observe_read_only_prints_state_and_never_fetches_waveform(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + code, out, _ = _run(["scope", "observe", "--channel", "1", "--config", str(config)]) + + assert code == 0 + assert "read_only=True" in out + assert "mutates_instrument=False" in out + assert "ch1 coupling=DC" in out + assert "waveform" not in out + assert _FakeScopeService.instances[0].fetched_channels == [] + + +def test_scope_observe_fetch_waveform_evaluates_expectations_and_recommends(): + with TemporaryDirectory() as tmp: + root = Path(tmp) + config = _write_config(root) + expectation = root / "expect.toml" + expectation.write_text("[channels.1]\nfrequency_hz = 1000.0\n", encoding="utf-8") + code, out, _ = _run( + [ + "scope", + "observe", + "--channel", + "1", + "--fetch-waveform", + "--expect", + str(expectation), + "--config", + str(config), + ] + ) + + assert code == 0 + assert _FakeScopeService.instances[0].fetched_channels == [1] + assert "mutates_instrument=True" in out + assert "ch1 expectation=ok" in out + assert "check=frequency_hz status=pass" in out + assert "recommendation" in out + + +def test_scope_observe_rejects_expect_file_without_fetch_waveform(): + with TemporaryDirectory() as tmp: + root = Path(tmp) + config = _write_config(root) + expectation = root / "expect.toml" + expectation.write_text("[channels.1]\nfrequency_hz = 1000.0\n", encoding="utf-8") + code, _, err = _run( + ["scope", "observe", "--expect", str(expectation), "--config", str(config)] + ) + + assert code != 0 + assert "--fetch-waveform" in err + assert _FakeScopeService.instances == [] + + +def test_scope_observe_rejects_invalid_expectation_before_any_instrument_io(): + with TemporaryDirectory() as tmp: + root = Path(tmp) + config = _write_config(root) + expectation = root / "expect.toml" + expectation.write_text("[channels.1]\nfrequncy_hz = 1000.0\n", encoding="utf-8") + code, _, err = _run( + [ + "scope", + "observe", + "--channel", + "1", + "--fetch-waveform", + "--expect", + str(expectation), + "--config", + str(config), + ] + ) + + assert code != 0 + assert "unknown expectation field" in err + # 校验失败必须发生在打开仪器之前 + assert _FakeScopeService.instances == [] + + +def test_scope_observe_rejects_expectation_for_unobserved_channel(): + with TemporaryDirectory() as tmp: + root = Path(tmp) + config = _write_config(root) + expectation = root / "expect.toml" + expectation.write_text("[channels.2]\nfrequency_hz = 1000.0\n", encoding="utf-8") + code, _, err = _run( + [ + "scope", + "observe", + "--channel", + "1", + "--fetch-waveform", + "--expect", + str(expectation), + "--config", + str(config), + ] + ) + + assert code != 0 + assert "expectation channels must be observed channels" in err + + +def test_json_scope_observe_wraps_result_in_versioned_envelope(): + with TemporaryDirectory() as tmp: + config = _write_config(Path(tmp)) + code, out, _ = _run(["--json", "scope", "observe", "--channel", "1", "--config", str(config)]) + + assert code == 0 + envelope = json.loads(out) + assert envelope["schema"] == "wavebench.cli.result.v1" + assert envelope["result"]["read_only"] is True + assert envelope["result"]["observation"]["channels"] == [1] + + +@pytest.mark.parametrize("option", ["--target-cycles", "--target-vertical-divisions"]) +@pytest.mark.parametrize("value", ["0", "-1", "nan", "inf", "-inf"]) +@pytest.mark.parametrize("fetch_waveform", [False, True]) +def test_scope_observe_rejects_invalid_targets_before_loading_config(option, value, fetch_waveform): + argv = ["scope", "observe", f"{option}={value}"] + if fetch_waveform: + argv.append("--fetch-waveform") + with patch("wavebench.services.agent_observe.load_config") as load: + code, _, err = _run(argv) + + assert code != 0 + assert option.removeprefix("--").replace("-", "_") in err + load.assert_not_called() + assert _FakeScopeService.instances == [] + + +@pytest.mark.parametrize("target", ["target_cycles", "target_vertical_divisions"]) +@pytest.mark.parametrize("value", [True, "1", object()]) +def test_scope_observe_validates_target_types_before_observation(target, value): + args = build_parser().parse_args(["scope", "observe", "--fetch-waveform"]) + setattr(args, target, value) + _FakeScopeService.instances = [] + with ( + patch("wavebench.services.agent_observe.load_config") as load, + patch("wavebench.services.agent_observe.ScopeService", _FakeScopeService), + pytest.raises(ConfigError, match=target), + ): + _run_scope_observe(args) + + load.assert_not_called() + assert _FakeScopeService.instances == [] diff --git a/tests/test_waveform_expectations.py b/tests/test_waveform_expectations.py new file mode 100644 index 00000000..60889f4f --- /dev/null +++ b/tests/test_waveform_expectations.py @@ -0,0 +1,227 @@ +import numpy as np +import pytest + +from wavebench.data.expectations import ( + estimate_triangle_symmetry_percent, + evaluate_waveform_expectation, + expectation_summary, + validate_expectation, +) +from wavebench.errors import ConfigError +from wavebench.instruments.models import WaveformData, WaveformHeader + + +def _waveform(channel: int, times: np.ndarray, values: np.ndarray) -> WaveformData: + return WaveformData( + channel=channel, + header=WaveformHeader(x_start=float(times[0]), x_stop=float(times[-1]), points=int(times.size)), + voltages_v=values, + ) + + +def test_square_wave_expectation_passes_frequency_vpp_mean_and_duty(): + times = np.linspace(0.0, 0.009999, 10_000) + values = np.where((times * 1000.0) % 1.0 < 0.5, 0.5, -0.5) + waveform = _waveform(1, times, values) + + result = evaluate_waveform_expectation( + waveform, + { + "label": "1k square", + "shape": "square", + "frequency_hz": 1000, + "frequency_tolerance_ratio": 0.02, + "vpp_v": 1.0, + "vpp_tolerance_ratio": 0.05, + "mean_v": 0.0, + "mean_tolerance_v": 0.02, + "duty_percent": 50, + "duty_tolerance": 0.02, + }, + ) + + assert result["status"] == "pass" + assert {check["metric"] for check in result["checks"]} == { + "frequency_hz", + "vpp_v", + "mean_v", + "duty_cycle", + } + + +def test_triangle_symmetry_expectation_passes_for_asymmetric_ramp(): + times = np.linspace(0.0, 0.0002, 5000) + period = 20e-6 + symmetry = 30.0 + phase = (times % period) / period + values = np.where( + phase < symmetry / 100.0, + -0.5 + phase / (symmetry / 100.0), + 0.5 - (phase - symmetry / 100.0) / (1.0 - symmetry / 100.0), + ) + values += 0.5 + waveform = _waveform(2, times, values) + + measured = estimate_triangle_symmetry_percent(waveform) + result = evaluate_waveform_expectation( + waveform, + { + "label": "50k triangle", + "shape": "triangle", + "frequency_hz": 50_000, + "vpp_v": 1.0, + "mean_v": 0.5, + "symmetry_percent": 30, + "symmetry_tolerance_percent": 3, + }, + ) + + assert measured is not None + assert abs(measured - 30.0) < 3.0 + assert result["status"] == "pass" + + +def test_expectation_warns_instead_of_failing_low_confidence_frequency(): + times = np.linspace(0.0, 0.0005, 200) + values = np.sin(2 * np.pi * 1000 * times) + waveform = _waveform(1, times, values) + + result = evaluate_waveform_expectation(waveform, {"frequency_hz": 1000}) + + assert result["status"] == "warn" + assert result["checks"][0]["status"] == "warn" + assert "low confidence" in result["checks"][0]["message"] + + +def test_expectation_summary_rolls_up_channel_statuses(): + summary = expectation_summary( + { + 1: {"status": "pass"}, + 2: {"status": "warn"}, + } + ) + + assert summary == {"status": "warn", "channels": {"1": "pass", "2": "warn"}} + + +@pytest.mark.parametrize(("statuses", "expected"), [ + ([], "skipped"), + (["skipped"], "skipped"), + (["unavailable"], "unavailable"), + (["unavailable", "skipped"], "unavailable"), + (["unavailable", "pass"], "partial"), + (["unavailable", "warn"], "partial"), + (["unavailable", "fail"], "fail"), +]) +def test_expectation_summary_accounts_for_unavailable_channels(statuses, expected): + results = {channel: {"status": status} for channel, status in enumerate(statuses, 1)} + + assert expectation_summary(results) == { + "status": expected, + "channels": {str(channel): result["status"] for channel, result in results.items()}, + } + + +def test_validate_expectation_rejects_unknown_field(): + # 拼错的字段名不能被静默忽略 + with pytest.raises(ConfigError, match="unknown expectation field"): + validate_expectation({"frequncy_hz": 1000}) + + +@pytest.mark.parametrize( + "expectation", + [ + {"frequency_hz": "1000"}, + {"frequency_hz": True}, + {"frequency_hz": float("nan")}, + {"frequency_hz": float("inf")}, + {"vpp_v": -1.0}, + {"vpp_v": 0.0}, + {"frequency_tolerance_ratio": -0.1}, + {"duty_cycle": 1.5}, + {"duty_percent": 120.0}, + {"symmetry_percent": 150.0}, + {"symmetry_tolerance_percent": -1.0}, + {"label": ""}, + ], +) +def test_validate_expectation_rejects_invalid_values(expectation): + with pytest.raises(ConfigError): + validate_expectation(expectation) + + +def test_validate_expectation_rejects_conflicting_synonyms(): + with pytest.raises(ConfigError, match="must not set both"): + validate_expectation({"duty_cycle": 0.5, "duty_percent": 50.0}) + with pytest.raises(ConfigError, match="must not set both"): + validate_expectation({"mean_v": 0.0, "offset_v": 0.0}) + + +def test_expectation_without_checkable_metric_is_skipped_not_passed(): + times = np.linspace(0.0, 0.001, 100) + waveform = _waveform(1, times, np.sin(2 * np.pi * 1000 * times)) + + # 只有 label/shape 时没有任何可执行检查,必须显式 skipped 而不是 pass + result = evaluate_waveform_expectation(waveform, {"label": "sine", "shape": "sine"}) + + assert result["status"] == "skipped" + assert result["checks"] == [] + assert "no checkable metric" in result["message"] + + +def test_evaluate_waveform_expectation_rejects_typo_before_any_check(): + times = np.linspace(0.0, 0.001, 100) + waveform = _waveform(1, times, np.sin(2 * np.pi * 1000 * times)) + + with pytest.raises(ConfigError, match="unknown expectation field"): + evaluate_waveform_expectation(waveform, {"frequncy_hz": 1000}) + + +def _triangle(times: np.ndarray, *, period: float, symmetry_percent: float, vpp: float) -> np.ndarray: + phase = (times % period) / period + rising = symmetry_percent / 100.0 + values = np.where( + phase < rising, + phase / rising, + 1.0 - (phase - rising) / (1.0 - rising), + ) + return values * vpp - vpp / 2.0 + + +@pytest.mark.parametrize("symmetry", [10.0, 50.0, 90.0]) +def test_triangle_symmetry_is_robust_to_noise_quantization_and_overshoot(symmetry): + period = 20e-6 + times = np.linspace(0.0, 20 * period, 20_000, endpoint=False) + period = float(times[1] - times[0]) * 1000.0 + clean = _triangle(times, period=period, symmetry_percent=symmetry, vpp=2.0) + + rng = np.random.default_rng(20260928) + # 1 mV 噪声 + 1 mV 量化台阶:逐点差分符号会被噪声打乱 + noisy = np.round(clean + rng.normal(0.0, 1e-3, clean.size), 3) + # 5 mV 噪声 + 轻微过冲 + overshoot = clean + rng.normal(0.0, 5e-3, clean.size) + 0.02 * np.sin(2 * np.pi * 5 / period * times) + + for values in (noisy, overshoot): + measured = estimate_triangle_symmetry_percent( + _waveform(1, times, values), + expected_frequency_hz=1.0 / period, + ) + assert measured is not None + assert abs(measured - symmetry) < 3.0, (symmetry, measured) + + +def test_triangle_symmetry_uses_expected_frequency_to_ignore_short_glitches(): + period = 20e-6 + times = np.linspace(0.0, 20 * period, 20_000, endpoint=False) + period = float(times[1] - times[0]) * 1000.0 + values = _triangle(times, period=period, symmetry_percent=10.0, vpp=2.0) + # 在上升沿插入一个远窄于半周期的毛刺;期望频率约束应把它排除在极值序列之外 + values[times.size // 2 : times.size // 2 + 20] += 0.3 + + measured = estimate_triangle_symmetry_percent( + _waveform(1, times, values), + expected_frequency_hz=1.0 / period, + ) + + assert measured is not None + assert abs(measured - 10.0) < 3.0 diff --git a/tests/test_waveform_relationships.py b/tests/test_waveform_relationships.py new file mode 100644 index 00000000..21880cdd --- /dev/null +++ b/tests/test_waveform_relationships.py @@ -0,0 +1,163 @@ +import numpy as np +import pytest + +from wavebench.data.relationships import analyze_waveform_pair, analyze_waveform_relationships +from wavebench.instruments.models import WaveformData, WaveformHeader + + +def _waveform(channel: int, values: np.ndarray, *, stop: float = 0.009) -> WaveformData: + return WaveformData( + channel=channel, + header=WaveformHeader(x_start=0.0, x_stop=stop, points=int(values.size)), + voltages_v=values, + ) + + +def test_waveform_pair_reports_frequency_voltage_and_phase_for_related_signals(): + t = np.linspace(0.0, 0.009, 1000) + left = _waveform(1, np.sin(2 * np.pi * 1000 * t), stop=float(t[-1])) + right = _waveform(2, 0.5 * np.sin(2 * np.pi * 1000 * (t - 0.00025)) + 0.2, stop=float(t[-1])) + + relationship = analyze_waveform_pair(left, right) + + assert relationship["channels"] == [1, 2] + assert relationship["common_time"]["overlap"] is True + assert relationship["frequency"]["ratio_high_over_low"] == 1.0 + assert 0.45 < relationship["voltage"]["vpp_ratio_right_over_left"] < 0.55 + assert 0.19 < relationship["voltage"]["mean_delta_right_minus_left_v"] < 0.21 + assert relationship["correlation"]["max_abs_cross_correlation"] > 0.9 + assert relationship["intersections"]["mode"] == "finite" + assert relationship["intersections"]["count"] > 0 + assert relationship["phase_degrees_at_left_frequency"] is not None + + +def test_waveform_pair_skips_timing_analysis_when_not_same_acquisition(): + t = np.linspace(0.0, 0.009, 1000) + left = _waveform(1, np.sin(2 * np.pi * 1000 * t), stop=float(t[-1])) + right = _waveform(2, np.sin(2 * np.pi * 1000 * (t - 0.00025)), stop=float(t[-1])) + + relationship = analyze_waveform_pair(left, right, same_acquisition=False) + + assert relationship["common_time"]["same_acquisition"] is False + assert relationship["common_time"]["overlap"] is None + assert relationship["phase_degrees_at_left_frequency"] is None + # 跨采集的波形没有共同时间基准,相关性和交点必须整段跳过而不是给出看似精确的数字 + assert relationship["correlation"] == {"status": "skipped", "reason": "not_same_acquisition"} + assert relationship["intersections"] == {"status": "skipped", "reason": "not_same_acquisition"} + assert "not_same_acquisition_timing_relationships_skipped" in relationship["warnings"] + # 同步无关的量仍然保留 + assert relationship["frequency"]["ratio_high_over_low"] == 1.0 + + +def test_waveform_pair_reports_phase_lag_in_degrees(): + t = np.linspace(0.0, 0.004, 4000) + left = _waveform(1, np.sin(2 * np.pi * 1000 * t), stop=float(t[-1])) + + for expected_degrees in (0.0, 90.0, 180.0, 270.0): + right = _waveform( + 2, + np.sin(2 * np.pi * 1000 * t - np.deg2rad(expected_degrees)), + stop=float(t[-1]), + ) + + relationship = analyze_waveform_pair(left, right) + + assert relationship["phase_degrees_at_left_frequency"] == pytest.approx( + expected_degrees, abs=0.5 + ) + + +def test_waveform_pair_reports_180_degrees_for_inverted_signal(): + t = np.linspace(0.0, 0.004, 4000) + left = np.sin(2 * np.pi * 1000 * t) + + relationship = analyze_waveform_pair( + _waveform(1, left, stop=float(t[-1])), + _waveform(2, -left, stop=float(t[-1])), + ) + + # 用相关峰绝对值选 lag 会把它报成 0° + assert relationship["phase_degrees_at_left_frequency"] == pytest.approx(180.0, abs=0.5) + + +def test_waveform_pair_phase_rejects_dc_leakage_in_noninteger_cycle_window(): + t = np.linspace(0.0, 0.0045, 4501) + left = _waveform(1, np.sin(2 * np.pi * 1000 * t), stop=float(t[-1])) + right = _waveform( + 2, np.sin(2 * np.pi * 1000 * t - np.pi / 2) + 5.0, stop=float(t[-1]), + ) + + relationship = analyze_waveform_pair(left, right) + + assert relationship["frequency"]["left_hz"] == pytest.approx(1000.0, abs=0.1) + assert relationship["frequency"]["right_hz"] == pytest.approx(1000.0, abs=0.1) + assert relationship["phase_degrees_at_left_frequency"] == pytest.approx(90.0, abs=0.5) + + +def test_waveform_relationships_report_all_pairs_for_four_channels(): + t = np.linspace(0.0, 0.004, 500) + waveforms = { + channel: _waveform(channel, np.sin(2 * np.pi * 1000 * t + channel), stop=float(t[-1])) + for channel in range(1, 5) + } + + relationships = analyze_waveform_relationships(waveforms) + + assert len(relationships) == 6 + assert relationships[0]["channels"] == [1, 2] + assert relationships[-1]["channels"] == [3, 4] + + +def test_waveform_pair_warns_when_frequency_confidence_is_low(): + t = np.linspace(0.0, 0.0005, 100) + left = _waveform(1, np.sin(2 * np.pi * 1000 * t), stop=float(t[-1])) + right = _waveform(2, np.sin(2 * np.pi * 2000 * t), stop=float(t[-1])) + + relationship = analyze_waveform_pair(left, right) + + assert relationship["frequency"]["left_hz"] is None + assert any("frequency_low_confidence" in warning for warning in relationship["warnings"]) + + +def test_waveform_pair_reports_intersection_points(): + t = np.linspace(0.0, 1.0, 1001) + left = _waveform(1, t - 0.25, stop=float(t[-1])) + right = _waveform(2, np.zeros_like(t), stop=float(t[-1])) + + relationship = analyze_waveform_pair(left, right) + + intersections = relationship["intersections"] + assert intersections["mode"] == "finite" + assert intersections["count"] == 1 + assert intersections["returned"] == 1 + assert intersections["truncated"] is False + assert intersections["points"][0]["time_s"] == 0.25 + assert intersections["points"][0]["voltage_v"] == 0.0 + assert intersections["points"][0]["direction"] == "left_minus_right_rising" + + +def test_waveform_pair_can_truncate_many_intersections(): + t = np.linspace(0.0, 0.01, 2000) + left = _waveform(1, np.sin(2 * np.pi * 1000 * t), stop=float(t[-1])) + right = _waveform(2, np.zeros_like(t), stop=float(t[-1])) + + relationship = analyze_waveform_pair(left, right, max_intersections=3) + + assert relationship["intersections"]["count"] > 3 + assert relationship["intersections"]["returned"] == 3 + assert relationship["intersections"]["truncated"] is True + assert "intersections_truncated" in relationship["warnings"] + + +def test_waveform_pair_marks_coincident_waveforms_as_unbounded_intersections(): + t = np.linspace(0.0, 0.001, 100) + values = np.sin(2 * np.pi * 1000 * t) + + relationship = analyze_waveform_pair( + _waveform(1, values, stop=float(t[-1])), + _waveform(2, values, stop=float(t[-1])), + ) + + assert relationship["intersections"]["mode"] == "coincident" + assert relationship["intersections"]["count"] is None + assert "waveforms_coincident_intersections_unbounded" in relationship["warnings"]