diff --git a/p3_vlm_orchestrator/08_policy_rollout.command b/p3_vlm_orchestrator/08_policy_rollout.command new file mode 100755 index 0000000..92c6c4e --- /dev/null +++ b/p3_vlm_orchestrator/08_policy_rollout.command @@ -0,0 +1,36 @@ +#!/bin/zsh + +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +REPO_ROOT="$(cd "$SCRIPT_DIR/.." && pwd)" +RUNTIME_ROOT="${RUNTIME_ROOT:-$REPO_ROOT/rebot_setup/vendor/rebot_lerobot}" + +if [[ -x "$RUNTIME_ROOT/reBotArm_control_py/.venv/bin/python" ]]; then + PYTHON_BIN="$RUNTIME_ROOT/reBotArm_control_py/.venv/bin/python" +elif [[ -x "$RUNTIME_ROOT/.venv/bin/python" ]]; then + PYTHON_BIN="$RUNTIME_ROOT/.venv/bin/python" +elif [[ -x "$REPO_ROOT/.venv/bin/python" ]]; then + PYTHON_BIN="$REPO_ROOT/.venv/bin/python" +else + PYTHON_BIN="$(command -v python3)" +fi + +typeset -a rollout_python_paths +rollout_python_paths=( + "$REPO_ROOT" + "$RUNTIME_ROOT/lerobot/src" + "$RUNTIME_ROOT/lerobot-robot-seeed-b601" +) +export RUNTIME_ROOT +export HF_LEROBOT_HOME="${HF_LEROBOT_HOME:-$RUNTIME_ROOT/lerobot-home}" +export PYTHONPATH="${(j/:/)rollout_python_paths}${PYTHONPATH:+:$PYTHONPATH}" + +print -- "Policy rollout safety reminder" +print -- "- Software q/x/Escape and signals are not a physical e-stop." +print -- "- Keep a dedicated operator on the physical e-stop/power cut for live mode." +print -- "- Stage offline, shadow, one-cycle live, then bounded multi-cycle live." +print -- "- Stopping/disconnecting releases follower torque; support the arm." +print -- "" + +exec "$PYTHON_BIN" -m p3_vlm_orchestrator.policy_rollout.cli "$@" diff --git a/p3_vlm_orchestrator/PERSON4_RUNBOOK.md b/p3_vlm_orchestrator/PERSON4_RUNBOOK.md new file mode 100644 index 0000000..e35f523 --- /dev/null +++ b/p3_vlm_orchestrator/PERSON4_RUNBOOK.md @@ -0,0 +1,213 @@ +# Person 4 policy rollout runbook + +This runbook covers Stage 1 only: `Pick up one can and place it in the taped sorting zone`. +It does not authorize multi-can clearing or unattended operation. + +## Locked policy contract + +- Actions and observations use seven physical follower-space joint positions, in + this exact order: `shoulder_pan, shoulder_lift, elbow_flex, wrist_flex, wrist_yaw, wrist_roll, gripper`. +- The image order is `front`, then `side`: `observation.images.front` is the + fixed Logitech overhead camera and `observation.images.side` is the Innomaker + wrist/claw camera. +- The adapter restores the checkpoint's saved LeRobot preprocessor and + postprocessor. Do not recreate normalization by hand. +- Each control cycle uses only the first predicted action, then obtains a fresh + observation and predicts again. + +Run every command from the DeskPartner repository root. Set the checkpoint and +dataset to absolute paths so the handoff is auditable: + +```zsh +CHECKPOINT=/absolute/path/to/checkpoint +DATASET=/absolute/path/to/finalized/lerobot-dataset +``` + +## Checkpoint handoff and inspection + +Person 3 must provide the complete immutable checkpoint directory, including +`model.safetensors`, `config.json`, `preprocessor_config.json`, +`postprocessor_config.json`, and `rebot_training_profile.json`. Record the +resolved checkpoint path and the `checkpoint_digest` printed by the shared +identity check in each trial manifest: + +```zsh +./p3_vlm_orchestrator/08_policy_rollout.command inspect --checkpoint "$CHECKPOINT" +``` + +Inspection must show the exact task, seven-joint order, `front,side` image +order, profile digest, and both saved processor configs. `inspect` validates +metadata and does not load policy weights or touch hardware. + +## Four mandatory gates + +Do not skip a gate. A fault or unexpected clamp returns the checkpoint to the +previous gate. + +### Gate A: offline recorded frames + +```zsh +./p3_vlm_orchestrator/08_policy_rollout.command offline \ + --checkpoint "$CHECKPOINT" --dataset "$DATASET" --episodes 2 --device cpu +``` + +Confirm finite `(10, 7)` action chunks on two distinct finalized episodes and +review the printed first-action deltas. This gate opens no serial port. + +### Gate B: shadow mode + +With the follower supported, the workspace clear, and a dedicated operator at +the physical e-stop/power cut, run prediction without sending actions: + +```zsh +./p3_vlm_orchestrator/08_policy_rollout.command shadow \ + --checkpoint "$CHECKPOINT" --cycles 20 --speed-scale 0.10 +``` + +Shadow mode still runs the calibrated workspace guard. It must send no action. + +### Gate C: one live cycle in an empty workspace + +Only after Gate B has no faults or clamps, use 10% speed: + +```zsh +./p3_vlm_orchestrator/08_policy_rollout.command live \ + --checkpoint "$CHECKPOINT" --cycles 1 --speed-scale 0.10 --live +``` + +Type both exact live acknowledgements when prompted. The independent e-stop +operator watches the arm, not the terminal. + +### Gate D: five live cycles in an empty workspace + +Only after a clean Gate C: + +```zsh +./p3_vlm_orchestrator/08_policy_rollout.command live \ + --checkpoint "$CHECKPOINT" --cycles 5 --speed-scale 0.10 --live +``` + +The first physical runs stay in an empty workspace at 10-20% speed. Do not put +a can in the workspace until all four gates have passed. + +## Held-out placement evaluation + +Reserve one fixed set of 10-15 distinct can placements that was not used for +training. Every checkpoint comparison must use the same placement IDs. A trial +is one placement after its final allowed outcome, not each attempt. + +For each placement, keep the physical e-stop operator present and run: + +```zsh +LOG_PATH="$PWD/runs/policy/checkpoint-a-held-out-01.jsonl" # must not exist yet +./p3_vlm_orchestrator/08_policy_rollout.command live \ + --checkpoint "$CHECKPOINT" --speed-scale 0.10 --live \ + --episode --retry-on-failure --log-path "$LOG_PATH" +``` + +The keyboard controls are `s` for operator success, `f` for operator failure, +and `q`, `x`, or Escape to stop. An operator failure may be retried once only +after the program has disconnected and a person has manually reset the can and +cleared the workspace. The retry prompts require the exact reset phrase +`I RESET THE CAN AND CLEARED THE WORKSPACE`, then both live acknowledgements +`I HAVE AN E-STOP OPERATOR` and `WORKSPACE IS EMPTY` again. There is no +automatic reset or motion. Safety faults, timeout, and stop are never retried. +Use a fresh, non-existing log path for every placement; never append a new +placement to an old JSONL. + +Create one read-only JSON manifest per checkpoint with exactly this envelope +and trial schema. Include 10-15 trial objects; the abbreviated example shows +one: + +```json +{ + "schema_version": 1, + "checkpoint": "/absolute/path/to/checkpoint", + "checkpoint_digest": "64-lowercase-hex-weight-digest", + "trials": [ + { + "checkpoint": "/absolute/path/to/checkpoint", + "checkpoint_digest": "64-lowercase-hex-weight-digest", + "placement_id": "held-out-01", + "attempts_used": 1, + "grasp_success": true, + "placement_success": true, + "terminal_reason": "operator_success", + "safety_faults": 0, + "clamps": 0, + "completion_s": 12.5, + "source_jsonl_paths": [ + "/absolute/path/to/runs/policy/checkpoint-a-held-out-01.jsonl" + ] + } + ] +} +``` + +Allowed terminal reasons are `operator_success`, `operator_failure`, `stopped`, +`timeout`, and `safety_fault`. Human reviewers label grasp and placement from +both camera views. `completion_s` is the sum of policy-running elapsed seconds +across the final trial's one or two attempts; exclude manual-reset downtime. +This same definition is used for every checkpoint. Sum clamp counts across both +attempts. `completion_s` must be a JSON number, not a quoted string. +`source_jsonl_paths` must identify absolute, existing, regular, non-symlink JSONL files. +The current retry loop appends both attempts to one file: List the shared JSONL path once, never duplicate it. +Reporting reads the files without modifying them and requires exactly the terminal or +`terminal_fallback` rows for attempts `1..attempts_used`. Their final reason, +total clamps, safety-fault count, and summed elapsed seconds must match the +manifest. A safety fault is a failed trial and is never retried. + +Generate the per-checkpoint report: + +```zsh +./p3_vlm_orchestrator/08_policy_rollout.command report \ + --manifest /absolute/path/to/checkpoint-a-trials.json \ + --output-name checkpoint-a-held-out +``` + +Compare two or more checkpoints: + +```zsh +./p3_vlm_orchestrator/08_policy_rollout.command compare \ + --manifest /absolute/path/to/checkpoint-a-trials.json \ + /absolute/path/to/checkpoint-b-trials.json \ + --output-name stage1-checkpoint-comparison +``` + +JSON and CSV are written under `runs/policy/reports/`; rollout audit JSONL is +written under `runs/policy/`. Existing reports are never overwritten. +Comparison ranks higher placement success first, then fewer safety faults, +then fewer clamps, then lower unrounded mean completion time. Reports round the +displayed mean to six decimal places only after ranking. Exact remaining ties +use checkpoint identity only for deterministic output. Training loss is never +a selection criterion. + +## Current fail-closed limitation + +The tracked `data/calibration/calibration.json` has null `plane_to_arm.A` and +`plane_to_arm.b` values, and the workspace corners in `config/workspace.yaml` +are also null. Shadow and live rollout therefore fail closed until a reviewed +physical workspace calibration for this rig is supplied. Never invent values, +bypass the guard, or weaken the check to make a run start. + +The adapter is generic across policies registered in the active LeRobot +runtime and always restores their saved processors. The bundled Python 3.11 / +LeRobot 0.4.4 runtime cannot execute MolmoAct2. It lacks the newer MolmoAct2 +plugin and dependencies. Gates A-D are therefore blocked for MolmoAct2 until +the team validates one combined Python 3.12 runtime that contains both the +MolmoAct2 policy stack and the ReBot hardware plugins, or implements and +validates a remote inference transport. Neither route exists in this repo +today. Generic loading does not mean missing policy dependencies work. + +## Stop, rollback, and shutdown + +Software stop is secondary to the physical e-stop. If motion is unsafe, the +dedicated operator immediately uses the physical e-stop or power cut; the +terminal operator also presses `q`, `x`, or Escape. Support the follower before +disconnect because the driver releases torque. After any stop or fault there +is no automatic home: do not command a return pose, restart, or continue with a +can present. Preserve the JSONL, inspect the terminal and fault rows, correct +the cause, clear the workspace, and restart at the last previously passed gate. + +No automated test performs physical motion. Automated verification uses pure +reporting data, fake adapters, and help/import checks only. diff --git a/p3_vlm_orchestrator/policy_rollout/__init__.py b/p3_vlm_orchestrator/policy_rollout/__init__.py new file mode 100644 index 0000000..18cfd74 --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/__init__.py @@ -0,0 +1,11 @@ +"""Hardware-independent policy rollout harness for Person 4.""" + +from .dummy_policy import HoldPositionPolicy, UnsafePolicy +from .runner import RolloutRunner, RunSummary + +__all__ = [ + "HoldPositionPolicy", + "RolloutRunner", + "RunSummary", + "UnsafePolicy", +] diff --git a/p3_vlm_orchestrator/policy_rollout/cli.py b/p3_vlm_orchestrator/policy_rollout/cli.py new file mode 100644 index 0000000..d069e9a --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/cli.py @@ -0,0 +1,1383 @@ +"""Gated command-line orchestration for offline and physical policy rollout. + +This module intentionally imports only the Python standard library. Optional +policy, NumPy, LeRobot, robot, SDK, and terminal helpers are imported only by +the subcommand handler that needs them. +""" + +from __future__ import annotations + +import argparse +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime, timezone +import hashlib +import json +import math +import os +from pathlib import Path +import re +import subprocess +import sys +import time +from typing import IO, Any + + +LIVE_OPERATOR_PHRASE = "I HAVE AN E-STOP OPERATOR" +EMPTY_WORKSPACE_PHRASE = "WORKSPACE IS EMPTY" +MIN_SPEED_SCALE = 0.10 +MAX_SPEED_SCALE = 0.20 +MAX_CLI_CYCLES = 20 +MANUAL_RESET_PHRASE = "I RESET THE CAN AND CLEARED THE WORKSPACE" +EXPECTED_IMAGE_ORDER = ( + "observation.images.front", + "observation.images.side", +) +EXPECTED_JOINT_NAMES = ( + "shoulder_pan", + "shoulder_lift", + "elbow_flex", + "wrist_flex", + "wrist_yaw", + "wrist_roll", + "gripper", +) +EXPECTED_COORDINATE_FRAME = ( + "follower_degrees_after_direction_limits_and_step_cap" +) +EXPECTED_CONTROL_MODE = "absolute joint pose" +PROFILE_RUNTIME_ENTRIES = ( + "follower", + "leader", + "follower_driver_contract", + "follower_base_implementation", + "follower_dm_implementation", + "leader_driver_contract", + "leader_implementation", +) +CAMERA_METADATA_KEYS = ( + "excluded_screen_index", + "excluded_screen_name", + "minimum_measured_fps", +) + + +HELP_EPILOG = """ +PHYSICAL SAFETY WARNING: software stop is not a physical e-stop. A dedicated +operator must hold the physical e-stop/power cut throughout every live cycle. + +Required staged rollout: + +# Gate A: offline only +python -m p3_vlm_orchestrator.policy_rollout.cli offline --checkpoint CHECKPOINT --dataset DATASET --episodes 2 + +# Gate B: hardware observation + prediction only, no send +python -m p3_vlm_orchestrator.policy_rollout.cli shadow --checkpoint CHECKPOINT --cycles 20 + +# Gate C: empty workspace, one cycle, 10% speed +python -m p3_vlm_orchestrator.policy_rollout.cli live --checkpoint CHECKPOINT --cycles 1 --speed-scale 0.10 --live + +# Gate D: only after Gate C has no faults/clamps, up to five cycles +python -m p3_vlm_orchestrator.policy_rollout.cli live --checkpoint CHECKPOINT --cycles 5 --speed-scale 0.10 --live + +Explicit held-out episode evaluation (Gate B-D behavior is unchanged unless +--episode is present): s = success, f = failure, q/x/Esc = stop. Each attempt +ends at the 30.0-second cap or 300 confirmed live actions. Stop and the physical +e-stop always take priority. --retry-on-failure permits at most one retry only +after an explicit manual reset acknowledgement; it never retries a safety +fault, timeout, or stop. +""" + + +@dataclass +class CliDependencies: + """Injectable boundaries used by tests and resolved lazily in production.""" + + stdout: IO[str] | None = None + stderr: IO[str] | None = None + input_fn: Callable[[str], str] | None = None + utc_now: Callable[[], datetime] | None = None + monotonic_clock: Callable[[], float] | None = None + checkpoint_loader: Callable[[Path], object] | None = None + offline_evaluator: Callable[..., object] | None = None + policy_factory: Callable[[object, str], object] | None = None + dummy_policy_factory: Callable[[], object] | None = None + safety_factory: Callable[..., object] | None = None + guard_factory: Callable[..., object] | None = None + robot_factory: Callable[..., object] | None = None + runner_factory: Callable[..., object] | None = None + keyboard_stop_factory: Callable[[], object] | None = None + serial_port_is_free: Callable[[str], bool] | None = None + repo_root: Path | None = None + + def output(self) -> IO[str]: + return self.stdout if self.stdout is not None else sys.stdout + + def errors(self) -> IO[str]: + return self.stderr if self.stderr is not None else sys.stderr + + def read_input(self) -> Callable[[str], str]: + return self.input_fn if self.input_fn is not None else input + + def now_utc(self) -> datetime: + now = self.utc_now() if self.utc_now is not None else datetime.now(timezone.utc) + if not isinstance(now, datetime) or now.tzinfo is None: + raise ValueError("UTC clock must return a timezone-aware datetime") + return now.astimezone(timezone.utc) + + def monotonic(self) -> Callable[[], float]: + return self.monotonic_clock or time.monotonic + + def repository_root(self) -> Path: + return self.repo_root if self.repo_root is not None else _repo_root() + + +@dataclass(frozen=True) +class _TerminalSummary: + cycles_completed: int + actions_attempted: int + actions_confirmed: int + terminal_reason: str + primary_fault_reason: str | None + cleanup_fault_reason: str | None + audit_fault_reason: str | None + attempt: int = 1 + elapsed_seconds: float = 0.0 + clamp_count: int = 0 + + +class PreflightRobotAdapter: + """Connect once, validate one discarded preflight, then delegate fresh reads.""" + + def __init__( + self, + *, + robot: object, + expected_task: str, + profile_snapshot: Mapping[str, Any], + ) -> None: + self.robot = robot + self.expected_task = expected_task + self.profile_snapshot = profile_snapshot + self._connect_attempted = False + self._connected = False + self._disconnected = False + self.cleanup_fault_reason: str | None = None + + def connect(self) -> None: + if self._connect_attempted: + raise RuntimeError("Preflight robot adapter cannot connect twice") + self._connect_attempted = True + try: + self.robot.connect() + self._connected = True + observation = self.robot.observe() + _validate_preflight_observation( + observation, + expected_task=self.expected_task, + profile_snapshot=self.profile_snapshot, + ) + except Exception: + try: + self._disconnect_once() + except Exception: + pass + raise + + def observe(self) -> object: + if not self._connected or self._disconnected: + raise RuntimeError("Preflight robot adapter is not connected") + return self.robot.observe() + + def send_action(self, action_deg: object) -> object: + if not self._connected or self._disconnected: + raise RuntimeError("Preflight robot adapter is not connected") + return self.robot.send_action(action_deg) + + def disconnect(self) -> None: + self._disconnect_once() + + def _disconnect_once(self) -> None: + if self._disconnected or not self._connect_attempted: + return + self._disconnected = True + self._connected = False + try: + self.robot.disconnect() + except Exception as exc: + self.cleanup_fault_reason = _exception_text(exc) + raise + + +def canonical_profile_digest(profile: object) -> str: + """Return the same canonical SHA-256 used by checkpoint sidecars.""" + + payload = json.dumps( + profile, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + ).encode("utf-8") + return hashlib.sha256(payload).hexdigest() + + +def default_serial_port_is_free( + port: str, + *, + runner: Callable[..., object] = subprocess.run, +) -> bool: + """Use ``lsof `` and fail closed on owners or checker errors.""" + + if not isinstance(port, str) or not port: + return False + try: + result = runner( + ["lsof", port], + capture_output=True, + text=True, + check=False, + ) + returncode = getattr(result, "returncode", None) + stdout = getattr(result, "stdout", "") + stderr = getattr(result, "stderr", "") + except Exception: + return False + if not isinstance(stdout, str) or not isinstance(stderr, str): + return False + if stdout.strip() or stderr.strip(): + return False + # lsof returns one when it found no matching open file. Every other status + # is either an owner (zero) or a checker failure (greater than one). + return returncode == 1 + + +def _repo_root() -> Path: + return Path(__file__).resolve().parents[2] + + +def _cycles(value: str) -> int: + try: + result = int(value) + except ValueError as exc: + raise argparse.ArgumentTypeError("cycles must be an integer") from exc + if not 1 <= result <= MAX_CLI_CYCLES: + raise argparse.ArgumentTypeError( + f"cycles must be within [1, {MAX_CLI_CYCLES}]" + ) + return result + + +def _episodes(value: str) -> int: + try: + result = int(value) + except ValueError as exc: + raise argparse.ArgumentTypeError("episodes must be an integer") from exc + if result <= 0: + raise argparse.ArgumentTypeError("episodes must be positive") + return result + + +def _add_hardware_arguments(parser: argparse.ArgumentParser) -> None: + root = _repo_root() + runtime_default = Path( + os.environ.get( + "RUNTIME_ROOT", + root / "rebot_setup" / "vendor" / "rebot_lerobot", + ) + ) + parser.add_argument("--cycles", type=_cycles) + parser.add_argument("--device", default="cpu") + parser.add_argument("--speed-scale", type=float, default=MIN_SPEED_SCALE) + parser.add_argument("--runtime-root", type=Path, default=runtime_default) + parser.add_argument("--arm-config", type=Path, default=root / "config" / "arm.yaml") + parser.add_argument( + "--workspace-config", + type=Path, + default=root / "config" / "workspace.yaml", + ) + parser.add_argument( + "--workspace-calibration", + type=Path, + default=root / "data" / "calibration" / "calibration.json", + ) + parser.add_argument("--log-path", type=Path, default=None) + parser.add_argument( + "--episode", + action="store_true", + help="Run one 30-second, human-verdict held-out evaluation attempt.", + ) + parser.add_argument( + "--retry-on-failure", + action="store_true", + help="After manual reset acknowledgement, retry operator failure once.", + ) + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="Safety-gated ReBot learned-policy rollout.", + epilog=HELP_EPILOG, + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + subparsers = parser.add_subparsers(dest="command", required=True) + + inspect_parser = subparsers.add_parser( + "inspect", + help="Validate and describe a checkpoint without loading policy weights.", + ) + inspect_parser.add_argument("--checkpoint", required=True, type=Path) + + offline_parser = subparsers.add_parser( + "offline", + help="Gate A: evaluate a checkpoint over finalized recorded episodes.", + ) + offline_parser.add_argument("--checkpoint", required=True, type=Path) + offline_parser.add_argument("--dataset", required=True, type=Path) + offline_parser.add_argument("--episodes", type=_episodes, default=2) + offline_parser.add_argument("--device", default="cpu") + + report_parser = subparsers.add_parser( + "report", + help="Validate one held-out trial manifest and write offline reports.", + ) + report_parser.add_argument("--manifest", required=True, type=Path) + report_parser.add_argument("--output-name", required=True) + + compare_parser = subparsers.add_parser( + "compare", + help="Rank checkpoints evaluated on the same held-out placements.", + ) + compare_parser.add_argument( + "--manifest", + required=True, + type=Path, + nargs="+", + help="Two or more strict final-trial JSON manifests.", + ) + compare_parser.add_argument("--output-name", required=True) + + shadow_parser = subparsers.add_parser( + "shadow", + help="Gate B: observe and predict without sending an action.", + ) + shadow_source = shadow_parser.add_mutually_exclusive_group(required=True) + shadow_source.add_argument("--checkpoint", type=Path) + shadow_source.add_argument("--dummy-hold", action="store_true") + shadow_parser.add_argument("--profile", type=Path) + _add_hardware_arguments(shadow_parser) + shadow_parser.set_defaults(cycles=20) + + live_parser = subparsers.add_parser( + "live", + help="Gates C-D: bounded physical action after every safety gate passes.", + ) + live_parser.add_argument("--checkpoint", required=True, type=Path) + live_parser.add_argument( + "--live", + dest="live_gate", + action="store_true", + help="Required acknowledgement that this subcommand may send motion.", + ) + _add_hardware_arguments(live_parser) + live_parser.set_defaults(cycles=1) + return parser + + +def main( + argv: Sequence[str] | None = None, + *, + dependencies: CliDependencies | None = None, +) -> int: + deps = dependencies or CliDependencies() + args = build_parser().parse_args(argv) + try: + if args.command == "inspect": + return _run_inspect(args, deps) + if args.command == "offline": + return _run_offline(args, deps) + if args.command == "report": + return _run_report(args, deps, compare_checkpoints=False) + if args.command == "compare": + return _run_report(args, deps, compare_checkpoints=True) + if args.command == "shadow": + return _run_hardware(args, deps, mode="shadow") + if args.command == "live": + return _run_hardware(args, deps, mode="live") + raise ValueError(f"Unsupported rollout command: {args.command}") + except Exception as exc: + print(f"ERROR: {_exception_text(exc)}", file=deps.errors()) + return 2 + + +def _checkpoint_loader(deps: CliDependencies) -> Callable[[Path], object]: + if deps.checkpoint_loader is not None: + return deps.checkpoint_loader + from rebot_operator_kit.rollout.checkpoint import CheckpointBundle + + return CheckpointBundle.load + + +def _run_inspect(args: argparse.Namespace, deps: CliDependencies) -> int: + bundle = _checkpoint_loader(deps)(args.checkpoint) + from p3_vlm_orchestrator.policy_rollout.evaluation import checkpoint_identity + + identity = checkpoint_identity(bundle.path) + config = _load_json_object(Path(bundle.path) / "config.json", "policy config") + policy_type = config.get("type", config.get("policy_type", "unknown")) + processor_artifacts = _processor_artifact_report(Path(bundle.path)) + coordinates = bundle.profile_snapshot.get("coordinate_contract", {}) + print(f"checkpoint={bundle.path}", file=deps.output()) + print(f"locked_task={bundle.task}", file=deps.output()) + print(f"policy_type={policy_type}", file=deps.output()) + print( + "policy_config=" + + json.dumps(config, sort_keys=True, separators=(",", ":")), + file=deps.output(), + ) + print( + "action_contract=" + f"dimension={bundle.action_dimension} chunk_size={bundle.chunk_size} " + f"action_steps={bundle.action_steps} " + f"frame={coordinates.get('frame')} control_mode={coordinates.get('control_mode')}", + file=deps.output(), + ) + joints = coordinates.get("joints", []) + joint_names = [ + joint.get("name") for joint in joints if isinstance(joint, Mapping) + ] + print("joint_order=" + ",".join(joint_names), file=deps.output()) + print("image_order=" + ",".join(bundle.image_order), file=deps.output()) + print( + "processor_artifacts=" + ",".join(processor_artifacts), + file=deps.output(), + ) + print(f"profile_digest={bundle.profile_digest}", file=deps.output()) + print(f"checkpoint_digest={identity.digest}", file=deps.output()) + return 0 + + +def _run_offline(args: argparse.Namespace, deps: CliDependencies) -> int: + bundle = _checkpoint_loader(deps)(args.checkpoint) + if deps.offline_evaluator is not None: + evaluator = deps.offline_evaluator + else: + from p3_vlm_orchestrator.policy_rollout.offline import evaluate_checkpoint + + evaluator = evaluate_checkpoint + results = evaluator( + bundle, + args.dataset, + episodes=args.episodes, + device=args.device, + output=deps.output(), + ) + if not results: + raise RuntimeError("No samples were found in the selected dataset episodes") + return 0 + + +def _run_report( + args: argparse.Namespace, + deps: CliDependencies, + *, + compare_checkpoints: bool, +) -> int: + """Load the standard-library-only reporting stack on explicit request.""" + + from p3_vlm_orchestrator.policy_rollout.evaluation import ( + compare, + load_trial_manifest, + write_reports, + ) + + manifest_paths = args.manifest if compare_checkpoints else [args.manifest] + manifests = [load_trial_manifest(path) for path in manifest_paths] + winner: str | None = None + if compare_checkpoints: + ranked = compare(manifests) + winner = str(ranked[0]["checkpoint"]) + paths = write_reports( + manifests, + output_name=args.output_name, + repo_root=deps.repository_root(), + ) + print(f"report_json={paths.json_path}", file=deps.output()) + print(f"report_csv={paths.csv_path}", file=deps.output()) + if winner is not None: + print(f"winner={winner}", file=deps.output()) + return 0 + + +def _run_hardware( + args: argparse.Namespace, + deps: CliDependencies, + *, + mode: str, +) -> int: + if mode == "live" and not args.live_gate: + raise ValueError("live rollout requires the explicit --live flag") + if args.retry_on_failure and not args.episode: + raise ValueError("--retry-on-failure requires explicit --episode mode") + speed_scale = _checked_speed_scale(args.speed_scale) + cycles = args.cycles + if not 1 <= cycles <= MAX_CLI_CYCLES: + raise ValueError(f"cycles must be within [1, {MAX_CLI_CYCLES}]") + log_path = _resolve_log_path( + args.log_path, + deps.now_utc(), + repo_root=deps.repository_root(), + ) + + bundle: object | None = None + if mode == "shadow" and args.dummy_hold: + if args.profile is None: + raise ValueError("shadow --dummy-hold requires --profile PATH") + profile_snapshot = _load_standalone_profile(args.profile) + profile_digest = canonical_profile_digest(profile_snapshot) + profile_authentication = "standalone-untrusted" + task = profile_snapshot["collection_defaults"]["task"] + checkpoint_path = None + checkpoint_digest = None + policy = _dummy_policy_factory(deps)() + else: + bundle = _checkpoint_loader(deps)(args.checkpoint) + profile_snapshot = bundle.profile_snapshot + profile_digest = bundle.profile_digest + profile_authentication = "checkpoint-sidecar-verified" + task = bundle.task + from p3_vlm_orchestrator.policy_rollout.evaluation import checkpoint_identity + + identity = checkpoint_identity(bundle.path) + checkpoint_path = str(identity.path) + checkpoint_digest = identity.digest + processor_artifacts = _processor_artifact_report(Path(bundle.path)) + print( + "processor_artifacts=" + ",".join(processor_artifacts), + file=deps.output(), + ) + policy = _policy_factory(deps)(bundle, args.device) + + if args.episode: + return _run_episode_hardware( + args, + deps, + mode=mode, + speed_scale=speed_scale, + log_path=log_path, + profile_snapshot=profile_snapshot, + profile_digest=profile_digest, + profile_authentication=profile_authentication, + checkpoint_path=checkpoint_path, + checkpoint_digest=checkpoint_digest, + task=task, + policy=policy, + ) + + safety = _safety_factory(deps)(profile_snapshot, mode=mode) + guard = _guard_factory(deps)( + arm_config_path=args.arm_config, + workspace_config_path=args.workspace_config, + calibration_path=args.workspace_calibration, + current_utc=deps.now_utc, + ) + robot = _robot_factory(deps)( + profile_snapshot=profile_snapshot, + runtime_root=args.runtime_root, + speed_scale=speed_scale, + monotonic_clock=deps.monotonic(), + ) + + follower_port = _selected_follower_port(robot) + serial_checker = deps.serial_port_is_free or default_serial_port_is_free + try: + serial_is_free = serial_checker(follower_port) + except Exception as exc: + raise RuntimeError(f"Follower serial ownership check failed: {exc}") from exc + if serial_is_free is not True: + raise RuntimeError("Follower serial port is owned or could not be checked") + + if mode == "live": + _require_live_phrases(deps.read_input()) + + _write_rollout_metadata( + log_path, + now=deps.now_utc(), + mode=mode, + task=task, + profile_digest=profile_digest, + profile_authentication=profile_authentication, + checkpoint_path=checkpoint_path, + checkpoint_digest=checkpoint_digest, + ) + print(f"jsonl_path={log_path}", file=deps.output()) + if mode == "shadow" and args.dummy_hold: + print( + f"standalone_profile_digest={profile_digest} (schema-validated, unauthenticated)", + file=deps.output(), + ) + + preflight_robot = PreflightRobotAdapter( + robot=robot, + expected_task=task, + profile_snapshot=profile_snapshot, + ) + keyboard = _keyboard_stop_factory(deps)() + summary: object + try: + with keyboard: + try: + runner = _runner_factory(deps)( + policy=policy, + robot=preflight_robot, + safety=safety, + mode=mode, + log_path=log_path, + monotonic_clock=deps.monotonic(), + stop_requested=keyboard.event, + action_guard=guard, + ) + except Exception as exc: + summary = _fault_summary( + primary=f"runner construction failed: {_exception_text(exc)}", + cleanup=preflight_robot.cleanup_fault_reason, + ) + else: + try: + summary = runner.run(cycles) + except Exception as exc: + summary = _fault_summary( + primary=f"runner execution failed: {_exception_text(exc)}", + cleanup=preflight_robot.cleanup_fault_reason, + ) + except Exception as exc: + summary = _fault_summary( + primary=f"keyboard stop setup failed: {_exception_text(exc)}", + cleanup=preflight_robot.cleanup_fault_reason, + ) + + summary = _with_cleanup_fault(summary, preflight_robot.cleanup_fault_reason) + + _print_summary(summary, deps.output()) + has_fault = any( + getattr(summary, field, None) + for field in ( + "primary_fault_reason", + "cleanup_fault_reason", + "audit_fault_reason", + ) + ) + return 1 if has_fault or summary.terminal_reason == "fault" else 0 + + +def _run_episode_hardware( + args: argparse.Namespace, + deps: CliDependencies, + *, + mode: str, + speed_scale: float, + log_path: Path, + profile_snapshot: Mapping[str, Any], + profile_digest: str, + profile_authentication: str, + checkpoint_path: str | None, + checkpoint_digest: str | None, + task: str, + policy: object, +) -> int: + """Run one episode and, only after manual reset, one fresh retry.""" + + attempt = 1 + metadata_written = False + while True: + reset_policy = getattr(policy, "reset", None) + if callable(reset_policy): + reset_policy() + safety = _safety_factory(deps)(profile_snapshot, mode=mode) + guard = _guard_factory(deps)( + arm_config_path=args.arm_config, + workspace_config_path=args.workspace_config, + calibration_path=args.workspace_calibration, + current_utc=deps.now_utc, + ) + robot = _robot_factory(deps)( + profile_snapshot=profile_snapshot, + runtime_root=args.runtime_root, + speed_scale=speed_scale, + monotonic_clock=deps.monotonic(), + ) + + follower_port = _selected_follower_port(robot) + serial_checker = deps.serial_port_is_free or default_serial_port_is_free + try: + serial_is_free = serial_checker(follower_port) + except Exception as exc: + raise RuntimeError( + f"Follower serial ownership check failed: {exc}" + ) from exc + if serial_is_free is not True: + raise RuntimeError("Follower serial port is owned or could not be checked") + + if mode == "live": + _require_live_phrases(deps.read_input()) + + if not metadata_written: + _write_rollout_metadata( + log_path, + now=deps.now_utc(), + mode=mode, + task=task, + profile_digest=profile_digest, + profile_authentication=profile_authentication, + checkpoint_path=checkpoint_path, + checkpoint_digest=checkpoint_digest, + ) + print(f"jsonl_path={log_path}", file=deps.output()) + if mode == "shadow" and args.dummy_hold: + print( + f"standalone_profile_digest={profile_digest} " + "(schema-validated, unauthenticated)", + file=deps.output(), + ) + metadata_written = True + + preflight_robot = PreflightRobotAdapter( + robot=robot, + expected_task=task, + profile_snapshot=profile_snapshot, + ) + keyboard = _keyboard_stop_factory(deps)() + summary: object + try: + with keyboard: + verdict_source = getattr(keyboard, "verdict", None) + if not callable(verdict_source): + raise RuntimeError( + "Episode keyboard control did not expose nonblocking verdicts" + ) + try: + runner = _runner_factory(deps)( + policy=policy, + robot=preflight_robot, + safety=safety, + mode=mode, + log_path=log_path, + monotonic_clock=deps.monotonic(), + stop_requested=keyboard.event, + action_guard=guard, + operator_verdict=verdict_source, + ) + except Exception as exc: + summary = _fault_summary( + primary=( + "runner construction failed: " + _exception_text(exc) + ), + cleanup=preflight_robot.cleanup_fault_reason, + episode=True, + attempt=attempt, + ) + else: + try: + summary = runner.run_episode(attempt=attempt) + except Exception as exc: + summary = _fault_summary( + primary=( + "runner execution failed: " + _exception_text(exc) + ), + cleanup=preflight_robot.cleanup_fault_reason, + episode=True, + attempt=attempt, + ) + except Exception as exc: + summary = _fault_summary( + primary="keyboard stop setup failed: " + _exception_text(exc), + cleanup=preflight_robot.cleanup_fault_reason, + episode=True, + attempt=attempt, + ) + + summary = _with_cleanup_fault( + summary, + preflight_robot.cleanup_fault_reason, + episode=True, + ) + _print_summary(summary, deps.output()) + has_fault = any( + getattr(summary, field, None) + for field in ( + "primary_fault_reason", + "cleanup_fault_reason", + "audit_fault_reason", + ) + ) + if has_fault or summary.terminal_reason == "safety_fault": + return 1 + if ( + summary.terminal_reason != "operator_failure" + or not args.retry_on_failure + or attempt >= 2 + ): + return 0 + + reset = deps.read_input()( + "After manual reset with no automatic motion, type exactly " + f'"{MANUAL_RESET_PHRASE}": ' + ) + if reset != MANUAL_RESET_PHRASE: + print( + "retry_skipped=manual_reset_not_acknowledged", + file=deps.output(), + ) + return 0 + attempt = 2 + + +def _policy_factory(deps: CliDependencies) -> Callable[[object, str], object]: + if deps.policy_factory is not None: + return deps.policy_factory + from p3_vlm_orchestrator.policy_rollout.lerobot_policy import ( + LeRobotPolicyAdapter, + ) + + return LeRobotPolicyAdapter.from_checkpoint + + +def _dummy_policy_factory(deps: CliDependencies) -> Callable[[], object]: + if deps.dummy_policy_factory is not None: + return deps.dummy_policy_factory + from p3_vlm_orchestrator.policy_rollout.dummy_policy import HoldPositionPolicy + + return HoldPositionPolicy + + +def _safety_factory(deps: CliDependencies) -> Callable[..., object]: + if deps.safety_factory is not None: + return deps.safety_factory + from rebot_operator_kit.rollout.safety import SafetyGovernor + + return SafetyGovernor.from_profile + + +def _guard_factory(deps: CliDependencies) -> Callable[..., object]: + if deps.guard_factory is not None: + return deps.guard_factory + from p3_vlm_orchestrator.policy_rollout.workspace_guard import ( + CalibratedWorkspaceGuard, + ) + + return CalibratedWorkspaceGuard.from_files + + +def _robot_factory(deps: CliDependencies) -> Callable[..., object]: + if deps.robot_factory is not None: + return deps.robot_factory + from p3_vlm_orchestrator.policy_rollout.rebot_robot import ReBotPolicyRobot + + return ReBotPolicyRobot.from_profile + + +def _runner_factory(deps: CliDependencies) -> Callable[..., object]: + if deps.runner_factory is not None: + return deps.runner_factory + from p3_vlm_orchestrator.policy_rollout.runner import RolloutRunner + + return RolloutRunner + + +def _keyboard_stop_factory(deps: CliDependencies) -> Callable[[], object]: + if deps.keyboard_stop_factory is not None: + return deps.keyboard_stop_factory + from p3_vlm_orchestrator.policy_rollout.keyboard_stop import KeyboardStop + + return KeyboardStop + + +def _checked_speed_scale(value: object) -> float: + if isinstance(value, bool): + raise ValueError("speed_scale must be within [0.10, 0.20]") + try: + speed = float(value) + except (TypeError, ValueError) as exc: + raise ValueError("speed_scale must be within [0.10, 0.20]") from exc + if not math.isfinite(speed) or not MIN_SPEED_SCALE <= speed <= MAX_SPEED_SCALE: + raise ValueError("speed_scale must be within [0.10, 0.20]") + return speed + + +def _selected_follower_port(robot: object) -> str: + for attribute in ("follower_port", "selected_follower_port"): + value = getattr(robot, attribute, None) + if isinstance(value, str) and value: + return value + follower = getattr(robot, "follower", None) + config = getattr(follower, "config", None) + value = getattr(config, "port", None) + if isinstance(value, str) and value: + return value + raise RuntimeError("Authenticated follower adapter did not expose its selected port") + + +def _require_live_phrases(input_fn: Callable[[str], str]) -> None: + operator = input_fn(f'Type exactly "{LIVE_OPERATOR_PHRASE}": ') + if operator != LIVE_OPERATOR_PHRASE: + raise RuntimeError("Physical e-stop operator confirmation was not received") + empty = input_fn(f'Type exactly "{EMPTY_WORKSPACE_PHRASE}": ') + if empty != EMPTY_WORKSPACE_PHRASE: + raise RuntimeError("Empty-workspace confirmation was not received") + + +def _validate_preflight_observation( + observation: object, + *, + expected_task: str, + profile_snapshot: Mapping[str, Any], +) -> None: + import numpy as np + + if getattr(observation, "task", None) != expected_task: + raise ValueError("Preflight observation task does not match the locked task") + try: + state = np.asarray(getattr(observation, "state_deg"), dtype=float) + except (TypeError, ValueError) as exc: + raise ValueError("Preflight state must contain seven finite values") from exc + if state.shape != (7,) or not np.isfinite(state).all(): + raise ValueError("Preflight state must contain seven finite values") + + cameras = _required_mapping(profile_snapshot, "camera_defaults", "profile") + front_contract = _required_mapping(cameras, "front", "camera contract") + side_contract = _required_mapping(cameras, "side", "camera contract") + _validate_preflight_image( + getattr(observation, "front", None), + label="front", + expected_height=front_contract.get("height"), + expected_width=front_contract.get("width"), + ) + _validate_preflight_image( + getattr(observation, "side", None), + label="side", + expected_height=side_contract.get("height"), + expected_width=side_contract.get("width"), + ) + + +def _validate_preflight_image( + image: object, + *, + label: str, + expected_height: object, + expected_width: object, +) -> None: + import numpy as np + + if ( + isinstance(expected_height, bool) + or not isinstance(expected_height, int) + or isinstance(expected_width, bool) + or not isinstance(expected_width, int) + ): + raise ValueError(f"Profile {label} camera dimensions are invalid") + array = np.asarray(image) + expected_shape = (expected_height, expected_width, 3) + if array.shape != expected_shape: + raise ValueError( + f"Preflight {label} image must be HWC RGB {expected_shape}; " + f"received {array.shape}" + ) + + +def _load_standalone_profile(path: Path) -> dict[str, Any]: + profile = _load_json_object(path, "standalone training profile") + _validate_standalone_profile(profile) + return profile + + +def _validate_standalone_profile(profile: Mapping[str, Any]) -> None: + if profile.get("schema_version") != 1: + raise ValueError("Standalone training profile schema_version must be 1") + profile_id = profile.get("profile_id") + if not isinstance(profile_id, str) or not re.fullmatch( + r"[a-z0-9][a-z0-9_-]{2,95}", profile_id + ): + raise ValueError("Standalone training profile_id is invalid") + version = profile.get("profile_version") + if isinstance(version, bool) or not isinstance(version, int) or version < 1: + raise ValueError("Standalone training profile_version must be positive") + + collection = _required_mapping(profile, "collection_defaults", "profile") + task = collection.get("task") + if not isinstance(task, str) or not task.strip(): + raise ValueError("Standalone training profile task must be nonempty") + if collection.get("motor_velocity") != 2000.0: + raise ValueError("Standalone training profile motor_velocity must equal 2000.0") + fps = _positive_integer(collection.get("fps"), "collection FPS") + gripper_force = _finite_number( + collection.get("gripper_force"), "collection gripper force" + ) + if not 0.0 <= gripper_force <= 1.0: + raise ValueError("Standalone collection gripper force must be in [0, 1]") + + coordinates = _required_mapping(profile, "coordinate_contract", "profile") + if coordinates.get("frame") != EXPECTED_COORDINATE_FRAME: + raise ValueError("Standalone training profile coordinate frame is invalid") + if coordinates.get("control_mode") != EXPECTED_CONTROL_MODE: + raise ValueError("Standalone training profile control mode is invalid") + if coordinates.get("action_dimension") != 7: + raise ValueError("Standalone training profile action dimension must be 7") + joints = coordinates.get("joints") + if not isinstance(joints, list) or len(joints) != 7: + raise ValueError("Standalone training profile must define exactly seven joints") + for index, (joint, expected_name) in enumerate( + zip(joints, EXPECTED_JOINT_NAMES, strict=True) + ): + if not isinstance(joint, Mapping): + raise ValueError(f"Standalone training profile joint {index} is invalid") + if joint.get("name") != expected_name or joint.get("feature") != f"{expected_name}.pos": + raise ValueError("Standalone training profile joint order/features are invalid") + _finite_nonzero(joint.get("leader_to_follower_scale"), f"joint {expected_name} scale") + limits = joint.get("soft_limit_degrees") + if not isinstance(limits, list) or len(limits) != 2: + raise ValueError(f"Standalone training profile joint {expected_name} limits are invalid") + low = _finite_number(limits[0], f"joint {expected_name} lower limit") + high = _finite_number(limits[1], f"joint {expected_name} upper limit") + if low >= high: + raise ValueError(f"Standalone training profile joint {expected_name} limits are invalid") + + training = _required_mapping(profile, "training_defaults", "profile") + if training.get("action_dimension") != 7: + raise ValueError("Standalone training policy action dimension must be 7") + _positive_integer(training.get("chunk_size"), "training chunk size") + _positive_integer(training.get("n_action_steps"), "training action steps") + if training.get("image_order") != list(EXPECTED_IMAGE_ORDER): + raise ValueError("Standalone training profile image order must be front then side") + if ( + training.get("state_normalization") != "quantile" + or training.get("action_normalization") != "quantile" + or training.get("normalize_gripper") is not True + ): + raise ValueError( + "Standalone training profile must lock quantile normalization and gripper normalization" + ) + + cameras = _required_mapping(profile, "camera_defaults", "profile") + extra = set(cameras) - {"front", "side"} - set(CAMERA_METADATA_KEYS) + if extra: + raise ValueError(f"Standalone camera profile has extra entries: {sorted(extra)}") + expected_cameras = { + "front": ("observation.images.front", 640, 480), + "side": ("observation.images.side", 1280, 720), + } + indices: list[int] = [] + for name, (recording_key, width, height) in expected_cameras.items(): + camera = _required_mapping(cameras, name, "camera contract") + if camera.get("recording_key") != recording_key: + raise ValueError(f"Standalone {name} camera recording key is invalid") + if camera.get("width") != width or camera.get("height") != height: + raise ValueError( + f"Standalone {name} camera must be exactly {width}x{height}" + ) + if camera.get("fps") != fps: + raise ValueError("Standalone camera FPS must match collection FPS") + indices.append(_nonnegative_integer(camera.get("index"), f"{name} camera index")) + if indices[0] == indices[1]: + raise ValueError("Standalone front and side camera indices must be distinct") + excluded_index = cameras.get("excluded_screen_index") + if excluded_index is not None: + excluded_index = _nonnegative_integer( + excluded_index, "excluded screen camera index" + ) + if excluded_index in indices: + raise ValueError("Standalone excluded screen camera must not be recorded") + minimum_fps = cameras.get("minimum_measured_fps") + if minimum_fps is not None: + minimum_fps = _positive_integer(minimum_fps, "minimum measured camera FPS") + if minimum_fps > fps: + raise ValueError("Standalone minimum measured camera FPS exceeds capture FPS") + + calibration = _required_mapping(profile, "calibration", "profile") + for name in PROFILE_RUNTIME_ENTRIES: + entry = _required_mapping(calibration, name, "profile calibration") + relative = entry.get("runtime_relative_path") + digest = entry.get("sha256") + if ( + not isinstance(relative, str) + or not relative + or Path(relative).is_absolute() + or ".." in Path(relative).parts + or not isinstance(digest, str) + or not re.fullmatch(r"[0-9a-f]{64}", digest) + ): + raise ValueError(f"Standalone profile calibration {name} is invalid") + follower = _required_mapping(calibration, "follower", "profile calibration") + if follower.get("type") != "seeed_b601_dm_follower" or follower.get("id") != "follower1": + raise ValueError("Standalone follower type/ID contract is invalid") + leader = _required_mapping(calibration, "leader", "profile calibration") + if leader.get("type") != "rebot_arm_102_leader" or leader.get("id") != "rebot_arm_102_leader": + raise ValueError("Standalone leader type/ID contract is invalid") + + hardware = _required_mapping(profile, "hardware_identity", "profile") + for identity_name in ("follower_usb", "leader_usb"): + usb_identity = _required_mapping( + hardware, identity_name, "hardware identity" + ) + for key in ("vid", "pid"): + value = usb_identity.get(key) + if ( + isinstance(value, bool) + or not isinstance(value, int) + or not 0 < value <= 0xFFFF + ): + raise ValueError(f"Standalone {identity_name} {key} is invalid") + + +def _required_mapping( + parent: Mapping[str, Any], key: str, label: str +) -> Mapping[str, Any]: + value = parent.get(key) + if not isinstance(value, Mapping): + raise ValueError(f"{label.capitalize()} {key} must be an object") + return value + + +def _finite_number(value: object, label: str) -> float: + if isinstance(value, bool): + raise ValueError(f"{label} must be finite") + try: + result = float(value) + except (TypeError, ValueError) as exc: + raise ValueError(f"{label} must be finite") from exc + if not math.isfinite(result): + raise ValueError(f"{label} must be finite") + return result + + +def _finite_nonzero(value: object, label: str) -> float: + result = _finite_number(value, label) + if result == 0.0: + raise ValueError(f"{label} must be nonzero") + return result + + +def _positive_integer(value: object, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise ValueError(f"{label} must be a positive integer") + return value + + +def _nonnegative_integer(value: object, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value < 0: + raise ValueError(f"{label} must be a nonnegative integer") + return value + + +def _load_json_object(path: Path, label: str) -> dict[str, Any]: + try: + value = json.loads( + Path(path).read_text(encoding="utf-8"), + parse_constant=_reject_json_constant, + ) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + raise ValueError(f"Cannot read {label} at {path}: {exc}") from exc + if not isinstance(value, dict): + raise ValueError(f"{label.capitalize()} must be an object") + return value + + +def _reject_json_constant(value: str) -> object: + raise ValueError(f"Nonfinite JSON constant is not permitted: {value}") + + +def _processor_artifact_report(checkpoint: Path) -> tuple[str, ...]: + root = checkpoint.resolve() + report: list[str] = [] + for filename in ("preprocessor_config.json", "postprocessor_config.json"): + document = _load_json_object(checkpoint / filename, "saved processor config") + report.append(filename) + steps = document.get("steps") + if not isinstance(steps, list): + raise ValueError(f"Saved processor config {filename} steps must be a list") + for index, step in enumerate(steps): + if not isinstance(step, Mapping): + raise ValueError(f"Saved processor config {filename} step {index} is invalid") + state_file = step.get("state_file") + if state_file is None: + continue + if not isinstance(state_file, str) or not state_file: + raise ValueError(f"Saved processor config {filename} state_file is invalid") + state_path = (checkpoint / state_file).resolve() + try: + state_path.relative_to(root) + except ValueError as exc: + raise ValueError("Saved processor state file leaves the checkpoint") from exc + if not state_path.is_file(): + raise ValueError(f"Saved processor state file is missing: {state_file}") + report.append(state_file) + return tuple(report) + + +def _resolve_log_path( + requested: Path | None, + now: datetime, + *, + repo_root: Path, +) -> Path: + """Resolve an audit path while confining it to the rollout log root.""" + + repository = Path(os.path.abspath(Path(repo_root).expanduser())) + allowed = repository / "runs" / "policy" + if requested is None: + timestamp = now.astimezone(timezone.utc).strftime("%Y%m%dT%H%M%SZ") + raw_target = allowed / f"{timestamp}.jsonl" + else: + supplied = Path(requested).expanduser() + if ".." in supplied.parts: + raise ValueError( + "Explicit log path must not traverse outside repo runs/policy/" + ) + raw_target = ( + supplied + if supplied.is_absolute() + else Path(os.path.abspath(supplied)) + ) + + raw_target = Path(os.path.abspath(raw_target)) + try: + relative = raw_target.relative_to(allowed) + except ValueError as exc: + raise ValueError( + "Explicit log path must be under repo runs/policy/; protected data, " + "models/checkpoints, config/calibration, env, and credential paths are forbidden" + ) from exc + if not relative.parts or raw_target.suffix != ".jsonl": + raise ValueError("Rollout log path under runs/policy/ must name a .jsonl file") + protected_names = (".env", "credential", "secret", "calibration", "checkpoint") + if any( + any(marker in part.lower() for marker in protected_names) + for part in relative.parts + ): + raise ValueError( + "Rollout log path under runs/policy/ cannot name protected env, " + "credential, calibration, or checkpoint files" + ) + + for directory in (repository / "runs", allowed): + if directory.is_symlink(): + raise ValueError("Rollout log path under runs/policy/ must not use symlinks") + candidate = allowed + for part in relative.parts: + candidate = candidate / part + if candidate.is_symlink(): + raise ValueError("Rollout log path under runs/policy/ must not use symlinks") + resolved_allowed = allowed.resolve(strict=False) + resolved_target = raw_target.resolve(strict=False) + try: + resolved_target.relative_to(resolved_allowed) + except ValueError as exc: + raise ValueError( + "Resolved rollout log path must remain under repo runs/policy/" + ) from exc + if resolved_target.exists() and not resolved_target.is_file(): + raise ValueError("Rollout log path under runs/policy/ must be a regular file") + return resolved_target + + +def _write_rollout_metadata( + path: Path, + *, + now: datetime, + mode: str, + task: str, + profile_digest: str, + profile_authentication: str, + checkpoint_path: str | None, + checkpoint_digest: str | None, +) -> None: + row = { + "event": "rollout_metadata", + "timestamp_utc": now.astimezone(timezone.utc).isoformat().replace("+00:00", "Z"), + "mode": mode, + "task": task, + "profile_digest": profile_digest, + "profile_authentication": profile_authentication, + "checkpoint": checkpoint_path, + "checkpoint_digest": checkpoint_digest, + } + try: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("x", encoding="utf-8") as handle: + handle.write(json.dumps(row, sort_keys=True, allow_nan=False) + "\n") + handle.flush() + except Exception as exc: + raise RuntimeError(f"Rollout metadata log could not be written: {exc}") from exc + + +def _fault_summary( + *, + primary: str, + cleanup: str | None = None, + audit: str | None = None, + episode: bool = False, + attempt: int = 1, +) -> _TerminalSummary: + return _TerminalSummary( + cycles_completed=0, + actions_attempted=0, + actions_confirmed=0, + terminal_reason="safety_fault" if episode else "fault", + primary_fault_reason=primary, + cleanup_fault_reason=cleanup, + audit_fault_reason=audit, + attempt=attempt, + ) + + +def _with_cleanup_fault( + summary: object, + cleanup: str | None, + *, + episode: bool = False, +) -> object: + if cleanup is None or getattr(summary, "cleanup_fault_reason", None) is not None: + return summary + return _TerminalSummary( + cycles_completed=getattr(summary, "cycles_completed"), + actions_attempted=getattr(summary, "actions_attempted"), + actions_confirmed=getattr(summary, "actions_confirmed"), + terminal_reason="safety_fault" if episode else "fault", + primary_fault_reason=getattr(summary, "primary_fault_reason", None), + cleanup_fault_reason=cleanup, + audit_fault_reason=getattr(summary, "audit_fault_reason", None), + attempt=getattr(summary, "attempt", 1), + elapsed_seconds=getattr(summary, "elapsed_seconds", 0.0), + clamp_count=getattr(summary, "clamp_count", 0), + ) + + +def _print_summary(summary: object, output: IO[str]) -> None: + print( + "rollout_summary " + f"cycles={getattr(summary, 'cycles_completed')} " + f"actions_attempted={getattr(summary, 'actions_attempted')} " + f"actions_confirmed={getattr(summary, 'actions_confirmed')} " + f"attempt={getattr(summary, 'attempt', 1)} " + f"elapsed_seconds={float(getattr(summary, 'elapsed_seconds', 0.0)):.3f} " + f"clamp_count={getattr(summary, 'clamp_count', 0)} " + f"terminal_reason={getattr(summary, 'terminal_reason')} " + f"primary_fault={_fault_text(getattr(summary, 'primary_fault_reason', None))} " + f"cleanup_fault={_fault_text(getattr(summary, 'cleanup_fault_reason', None))} " + f"audit_fault={_fault_text(getattr(summary, 'audit_fault_reason', None))}", + file=output, + ) + + +def _fault_text(value: object) -> str: + return "none" if value is None else str(value) + + +def _exception_text(exc: Exception) -> str: + try: + return str(exc) + except Exception: + return type(exc).__name__ + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/p3_vlm_orchestrator/policy_rollout/dummy_policy.py b/p3_vlm_orchestrator/policy_rollout/dummy_policy.py new file mode 100644 index 0000000..097f822 --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/dummy_policy.py @@ -0,0 +1,21 @@ +"""Small deterministic policies for exercising the rollout harness.""" + +from __future__ import annotations + +import numpy as np + +from rebot_operator_kit.rollout.contracts import RolloutObservation + + +class HoldPositionPolicy: + """Return a ten-step chunk that holds the observed follower pose.""" + + def predict(self, observation: RolloutObservation) -> np.ndarray: + return np.repeat(observation.state_deg[None, :], 10, axis=0) + + +class UnsafePolicy: + """Return intentionally invalid actions for fail-closed smoke tests.""" + + def predict(self, observation: RolloutObservation) -> np.ndarray: + return np.full((10, 7), np.nan) diff --git a/p3_vlm_orchestrator/policy_rollout/evaluation.py b/p3_vlm_orchestrator/policy_rollout/evaluation.py new file mode 100644 index 0000000..f642180 --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/evaluation.py @@ -0,0 +1,535 @@ +"""Hardware-free held-out rollout validation, aggregation, and reporting.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +import csv +from dataclasses import dataclass +import hashlib +from io import StringIO +import json +import math +import os +from pathlib import Path +import re +from typing import Any + + +SUMMARY_FIELDS = ( + "checkpoint", + "trials", + "grasp_successes", + "placement_successes", + "safety_faults", + "clamps", + "mean_completion_s", + "overall_success_rate", +) +MANIFEST_FIELDS = frozenset( + ("schema_version", "checkpoint", "checkpoint_digest", "trials") +) +TRIAL_FIELDS = frozenset( + ( + "checkpoint", + "checkpoint_digest", + "placement_id", + "attempts_used", + "grasp_success", + "placement_success", + "terminal_reason", + "safety_faults", + "clamps", + "completion_s", + "source_jsonl_paths", + ) +) +TERMINAL_REASONS = frozenset( + ("operator_success", "operator_failure", "stopped", "timeout", "safety_fault") +) +_OUTPUT_NAME = re.compile(r"[a-z0-9][a-z0-9_-]{0,63}") +_DIGEST = re.compile(r"[0-9a-f]{64}") +_PROTECTED_OUTPUT_MARKERS = ( + "credential", + "secret", + "calibration", + "dataset", + ".env", +) + + +@dataclass(frozen=True) +class Trial: + checkpoint: str + checkpoint_digest: str + placement_id: str + attempts_used: int + grasp_success: bool + placement_success: bool + terminal_reason: str + safety_faults: int + clamps: int + completion_s: float + source_jsonl_paths: tuple[str, ...] + + +@dataclass(frozen=True) +class TrialManifest: + source_path: Path + checkpoint: str + checkpoint_digest: str + trials: tuple[Trial, ...] + + +@dataclass(frozen=True) +class ReportPaths: + json_path: Path + csv_path: Path + + +@dataclass(frozen=True) +class CheckpointIdentity: + path: Path + digest: str + + +def checkpoint_identity(path: Path | str) -> CheckpointIdentity: + """Bind the accepted checkpoint layout to its resolved path and weights.""" + + checkpoint = Path(path).expanduser().resolve() + weights = checkpoint / "model.safetensors" + if weights.is_symlink() or not weights.is_file(): + raise ValueError("Checkpoint model.safetensors is missing or not a regular file") + digest = hashlib.sha256() + try: + with weights.open("rb") as handle: + for chunk in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(chunk) + except OSError as exc: + raise ValueError(f"Checkpoint model.safetensors cannot be read: {exc}") from exc + return CheckpointIdentity(path=checkpoint, digest=digest.hexdigest()) + + +@dataclass(frozen=True) +class _TerminalAudit: + attempt: int + terminal_reason: str + clamps: int + elapsed_seconds: float + + +def load_trial_manifest(path: Path | str) -> TrialManifest: + """Read and strictly validate one checkpoint's final held-out trials.""" + + source = Path(path) + try: + document = json.loads( + source.read_text(encoding="utf-8"), + parse_constant=_reject_json_constant, + ) + except (OSError, UnicodeError, json.JSONDecodeError, ValueError) as exc: + raise ValueError(f"Cannot read trial manifest at {source}: {exc}") from exc + if not isinstance(document, Mapping) or set(document) != MANIFEST_FIELDS: + raise ValueError("Trial manifest schema is not exact") + if document.get("schema_version") != 1: + raise ValueError("Trial manifest schema_version must be exactly 1") + + checkpoint = _nonempty_string(document.get("checkpoint"), "checkpoint") + digest = _checkpoint_digest(document.get("checkpoint_digest")) + rows = document.get("trials") + if not isinstance(rows, list) or not 10 <= len(rows) <= 15: + raise ValueError("Trial manifest must contain 10 to 15 final trials") + + trials = tuple( + _parse_trial(row, index=index, checkpoint=checkpoint, digest=digest) + for index, row in enumerate(rows) + ) + placement_ids = [trial.placement_id for trial in trials] + if len(set(placement_ids)) != len(placement_ids): + raise ValueError("Trials must use distinct placement IDs") + audit_sources = [ + source + for trial in trials + for source in trial.source_jsonl_paths + ] + if len(set(audit_sources)) != len(audit_sources): + raise ValueError("source JSONL path is reused across placements") + return TrialManifest( + source_path=source, + checkpoint=checkpoint, + checkpoint_digest=digest, + trials=trials, + ) + + +def summarize(manifest: TrialManifest) -> dict[str, object]: + """Aggregate one validated manifest using every final trial duration.""" + + trials = manifest.trials + count = len(trials) + placement_successes = sum(trial.placement_success for trial in trials) + return { + "checkpoint": manifest.checkpoint, + "trials": count, + "grasp_successes": sum(trial.grasp_success for trial in trials), + "placement_successes": placement_successes, + "safety_faults": sum(trial.safety_faults for trial in trials), + "clamps": sum(trial.clamps for trial in trials), + "mean_completion_s": round( + math.fsum(trial.completion_s for trial in trials) / count, 6 + ), + "overall_success_rate": round(placement_successes / count, 6), + } + + +def compare(manifests: Sequence[TrialManifest]) -> list[dict[str, object]]: + """Rank checkpoints on an identical held-out placement set.""" + + if len(manifests) < 2: + raise ValueError("Checkpoint comparison requires at least two manifests") + checkpoint_paths = [manifest.checkpoint for manifest in manifests] + if len(set(checkpoint_paths)) != len(checkpoint_paths): + raise ValueError("Checkpoint comparison requires distinct checkpoint identities") + checkpoint_digests = [manifest.checkpoint_digest for manifest in manifests] + if len(set(checkpoint_digests)) != len(checkpoint_digests): + raise ValueError("Checkpoint comparison requires distinct checkpoint digests") + expected_placements = frozenset( + trial.placement_id for trial in manifests[0].trials + ) + for manifest in manifests[1:]: + placements = frozenset(trial.placement_id for trial in manifest.trials) + if placements != expected_placements: + raise ValueError( + "Checkpoint comparison requires the same held-out placement IDs" + ) + + ranked = sorted( + manifests, + key=lambda manifest: ( + -sum(trial.placement_success for trial in manifest.trials), + sum(trial.safety_faults for trial in manifest.trials), + sum(trial.clamps for trial in manifest.trials), + math.fsum(trial.completion_s for trial in manifest.trials) + / len(manifest.trials), + manifest.checkpoint, + ), + ) + return [summarize(manifest) for manifest in ranked] + + +def write_reports( + manifests: Sequence[TrialManifest], + *, + output_name: str, + repo_root: Path | str, +) -> ReportPaths: + """Write deterministic JSON and CSV without overwriting prior reports.""" + + report_root = _report_root(repo_root, output_name) + json_path = report_root / f"{output_name}.json" + csv_path = report_root / f"{output_name}.csv" + for output in (json_path, csv_path): + if output.exists() or output.is_symlink(): + raise ValueError(f"Report output already exists: {output}") + + if len(manifests) == 1: + rows = [summarize(manifests[0])] + else: + rows = compare(manifests) + _validate_summary_rows(rows) + + json_text = json.dumps(rows, indent=2, ensure_ascii=False, allow_nan=False) + "\n" + csv_buffer = StringIO(newline="") + writer = csv.DictWriter( + csv_buffer, + fieldnames=SUMMARY_FIELDS, + extrasaction="raise", + lineterminator="\n", + ) + writer.writeheader() + writer.writerows(rows) + csv_text = csv_buffer.getvalue() + + created: list[Path] = [] + try: + for output, text in ((json_path, json_text), (csv_path, csv_text)): + with output.open("x", encoding="utf-8", newline="") as handle: + handle.write(text) + handle.flush() + created.append(output) + except Exception: + for output in created: + output.unlink(missing_ok=True) + raise + return ReportPaths(json_path=json_path.resolve(), csv_path=csv_path.resolve()) + + +def _parse_trial( + value: object, + *, + index: int, + checkpoint: str, + digest: str, +) -> Trial: + if not isinstance(value, Mapping) or set(value) != TRIAL_FIELDS: + raise ValueError(f"Trial {index} schema is not exact") + trial_checkpoint = _nonempty_string(value.get("checkpoint"), "trial checkpoint") + trial_digest = _checkpoint_digest(value.get("checkpoint_digest")) + if trial_checkpoint != checkpoint or trial_digest != digest: + raise ValueError(f"Trial {index} has mixed checkpoint identity or digest") + placement_id = _nonempty_string(value.get("placement_id"), "placement ID") + attempts_used = _integer(value.get("attempts_used"), "attempts_used") + if attempts_used not in (1, 2): + raise ValueError("attempts_used must be exactly 1 or 2") + grasp_success = _boolean(value.get("grasp_success"), "grasp_success") + placement_success = _boolean( + value.get("placement_success"), "placement_success" + ) + terminal_reason = _nonempty_string( + value.get("terminal_reason"), "terminal_reason" + ) + if terminal_reason not in TERMINAL_REASONS: + raise ValueError(f"Unrecognized terminal reason: {terminal_reason}") + safety_faults = _nonnegative_integer( + value.get("safety_faults"), "safety_faults" + ) + clamps = _nonnegative_integer(value.get("clamps"), "clamps") + completion_s = _nonnegative_finite(value.get("completion_s"), "completion_s") + source_paths = _source_paths(value.get("source_jsonl_paths")) + + if placement_success and not grasp_success: + raise ValueError("Placement success requires grasp success") + if terminal_reason == "operator_success" and not placement_success: + raise ValueError("Operator success requires placement success") + if placement_success and terminal_reason != "operator_success": + raise ValueError("Placement success requires operator_success terminal reason") + if terminal_reason == "safety_fault": + if grasp_success or placement_success: + raise ValueError("Safety-fault terminal outcomes cannot be successful") + if safety_faults < 1: + raise ValueError("Safety-fault terminal outcome must count a safety fault") + elif safety_faults != 0: + raise ValueError("Safety fault counts require a safety_fault terminal outcome") + + trial = Trial( + checkpoint=trial_checkpoint, + checkpoint_digest=trial_digest, + placement_id=placement_id, + attempts_used=attempts_used, + grasp_success=grasp_success, + placement_success=placement_success, + terminal_reason=terminal_reason, + safety_faults=safety_faults, + clamps=clamps, + completion_s=completion_s, + source_jsonl_paths=source_paths, + ) + _validate_trial_audits(trial) + return trial + + +def _report_root(repo_root: Path | str, output_name: str) -> Path: + if not isinstance(output_name, str) or not _OUTPUT_NAME.fullmatch(output_name): + raise ValueError("Report output name must be a lowercase slug") + if any(marker in output_name for marker in _PROTECTED_OUTPUT_MARKERS): + raise ValueError("Report output name cannot reference protected content") + repository = Path(os.path.abspath(Path(repo_root).expanduser())) + report_root = repository / "runs" / "policy" / "reports" + for path in ( + repository, + repository / "runs", + repository / "runs" / "policy", + report_root, + ): + if path.is_symlink(): + raise ValueError("Report output path must not use a symlink") + report_root.mkdir(parents=True, exist_ok=True) + if report_root.is_symlink(): + raise ValueError("Report output path must not use a symlink") + resolved_root = report_root.resolve(strict=True) + expected_root = (repository / "runs" / "policy" / "reports").resolve(strict=True) + if resolved_root != expected_root: + raise ValueError("Report output must remain under runs/policy/reports") + return resolved_root + + +def _source_paths(value: object) -> tuple[str, ...]: + if not isinstance(value, list) or not value: + raise ValueError("source JSONL paths must be a nonempty list") + paths = tuple(_nonempty_string(item, "source JSONL path") for item in value) + if len(set(paths)) != len(paths): + raise ValueError("source JSONL paths must be distinct") + for source in paths: + path = Path(source) + if not path.is_absolute(): + raise ValueError("source JSONL paths must be absolute") + if path.suffix != ".jsonl": + raise ValueError("source JSONL paths must name .jsonl files") + if path.is_symlink() or not path.is_file(): + raise ValueError( + "source JSONL paths must be existing regular non-symlink files" + ) + return paths + + +def _validate_trial_audits(trial: Trial) -> None: + terminal_rows: list[_TerminalAudit] = [] + for source in trial.source_jsonl_paths: + terminal_rows.extend( + _read_terminal_audits( + Path(source), + checkpoint=trial.checkpoint, + checkpoint_digest=trial.checkpoint_digest, + ) + ) + + attempts = sorted(row.attempt for row in terminal_rows) + expected_attempts = list(range(1, trial.attempts_used + 1)) + if attempts != expected_attempts: + raise ValueError( + "source JSONL terminal attempts must be exactly 1..attempts_used" + ) + ordered = sorted(terminal_rows, key=lambda row: row.attempt) + if any(row.terminal_reason != "operator_failure" for row in ordered[:-1]): + raise ValueError( + "source JSONL non-final attempts must end in operator_failure" + ) + if ordered[-1].terminal_reason != trial.terminal_reason: + raise ValueError( + "source JSONL final terminal reason does not match the manifest" + ) + + audit_clamps = sum(row.clamps for row in ordered) + if audit_clamps != trial.clamps: + raise ValueError("source JSONL clamp count does not match the manifest") + audit_safety_faults = sum( + row.terminal_reason == "safety_fault" for row in ordered + ) + if audit_safety_faults != trial.safety_faults: + raise ValueError( + "source JSONL safety fault count does not match the manifest" + ) + audit_completion_s = math.fsum(row.elapsed_seconds for row in ordered) + if not math.isclose( + audit_completion_s, + trial.completion_s, + rel_tol=1e-9, + abs_tol=1e-9, + ): + raise ValueError( + "source JSONL completion seconds do not match the manifest" + ) + + +def _read_terminal_audits( + path: Path, + *, + checkpoint: str, + checkpoint_digest: str, +) -> list[_TerminalAudit]: + try: + lines = path.read_text(encoding="utf-8").splitlines() + except (OSError, UnicodeError) as exc: + raise ValueError(f"Cannot read source JSONL at {path}: {exc}") from exc + terminal_rows: list[_TerminalAudit] = [] + metadata_rows: list[Mapping[str, object]] = [] + for line_number, line in enumerate(lines, start=1): + if not line.strip(): + raise ValueError(f"Source JSONL {path}:{line_number} is blank") + try: + row = json.loads(line, parse_constant=_reject_json_constant) + except (json.JSONDecodeError, ValueError) as exc: + raise ValueError( + f"Source JSONL {path}:{line_number} is invalid: {exc}" + ) from exc + if not isinstance(row, Mapping): + raise ValueError( + f"Source JSONL {path}:{line_number} must contain an object" + ) + if row.get("event") == "rollout_metadata": + metadata_rows.append(row) + continue + if row.get("event") not in ("terminal", "terminal_fallback"): + continue + attempt = _integer(row.get("attempt"), "source JSONL terminal attempt") + if attempt < 1: + raise ValueError("source JSONL terminal attempt must be positive") + terminal_reason = _nonempty_string( + row.get("terminal_reason"), + "source JSONL terminal reason", + ) + if terminal_reason not in TERMINAL_REASONS: + raise ValueError("source JSONL terminal reason is unrecognized") + terminal_rows.append( + _TerminalAudit( + attempt=attempt, + terminal_reason=terminal_reason, + clamps=_nonnegative_integer( + row.get("clamp_count"), + "source JSONL clamp count", + ), + elapsed_seconds=_nonnegative_finite( + row.get("elapsed_seconds"), + "source JSONL elapsed seconds", + ), + ) + ) + if len(metadata_rows) != 1: + raise ValueError("source JSONL must contain exactly one rollout_metadata row") + metadata = metadata_rows[0] + if metadata.get("profile_authentication") != "checkpoint-sidecar-verified": + raise ValueError("source JSONL does not contain an authenticated checkpoint") + if metadata.get("checkpoint") != checkpoint: + raise ValueError("source JSONL checkpoint does not match the manifest") + if metadata.get("checkpoint_digest") != checkpoint_digest: + raise ValueError("source JSONL checkpoint digest does not match the manifest") + return terminal_rows + + +def _validate_summary_rows(rows: Sequence[Mapping[str, object]]) -> None: + for row in rows: + if tuple(row) != SUMMARY_FIELDS: + raise ValueError("Summary schema is not exact") + + +def _checkpoint_digest(value: object) -> str: + if not isinstance(value, str) or not _DIGEST.fullmatch(value): + raise ValueError("checkpoint_digest must be a lowercase SHA-256 digest") + return value + + +def _nonempty_string(value: object, label: str) -> str: + if not isinstance(value, str) or not value or value != value.strip(): + raise ValueError(f"{label} must be a nonempty trimmed string") + return value + + +def _integer(value: object, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise ValueError(f"{label} must be an integer") + return value + + +def _nonnegative_integer(value: object, label: str) -> int: + result = _integer(value, label) + if result < 0: + raise ValueError(f"{label} must be nonnegative") + return result + + +def _nonnegative_finite(value: object, label: str) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError(f"{label} must be a real JSON number") + result = float(value) + if not math.isfinite(result) or result < 0.0: + raise ValueError(f"{label} must be finite and nonnegative") + return result + + +def _boolean(value: object, label: str) -> bool: + if not isinstance(value, bool): + raise ValueError(f"{label} must be boolean") + return value + + +def _reject_json_constant(value: str) -> object: + raise ValueError(f"Nonfinite JSON constant is not permitted: {value}") diff --git a/p3_vlm_orchestrator/policy_rollout/keyboard_stop.py b/p3_vlm_orchestrator/policy_rollout/keyboard_stop.py new file mode 100644 index 0000000..c8b5df0 --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/keyboard_stop.py @@ -0,0 +1,208 @@ +"""Permission-free terminal and signal stop handling for policy rollout.""" + +from __future__ import annotations + +from collections.abc import Callable +import select +import signal +import sys +import termios +import threading +import tty +from typing import IO, Any + + +STOP_KEYS = frozenset(("q", "x", "\x1b")) +SUCCESS_KEYS = frozenset(("s",)) +FAILURE_KEYS = frozenset(("f",)) +POLL_INTERVAL_S = 0.10 + + +class KeyboardStop: + """Route terminal keys and process signals into one shared stop event. + + Construction and import are side-effect free. Signal handlers, terminal mode, + and the daemon reader exist only while the context manager is active. + """ + + def __init__( + self, + *, + event: threading.Event | None = None, + stdin: IO[str] | None = None, + warning_stream: IO[str] | None = None, + signal_api: Any = None, + termios_api: Any = None, + tty_api: Any = None, + select_fn: Callable[..., tuple[list[object], list[object], list[object]]] | None = None, + thread_factory: Callable[..., object] | None = None, + is_main_thread: Callable[[], bool] | None = None, + ) -> None: + self.event = event or threading.Event() + self.stdin = stdin if stdin is not None else sys.stdin + self.warning_stream = ( + warning_stream if warning_stream is not None else sys.stderr + ) + self.signal_api = signal_api if signal_api is not None else signal + self.termios_api = termios_api if termios_api is not None else termios + self.tty_api = tty_api if tty_api is not None else tty + self.select_fn = select_fn if select_fn is not None else select.select + self.thread_factory = ( + thread_factory if thread_factory is not None else threading.Thread + ) + self.is_main_thread = is_main_thread or ( + lambda: threading.current_thread() is threading.main_thread() + ) + self._prior_handlers: dict[object, object] = {} + self._terminal_fd: int | None = None + self._terminal_state: object | None = None + self._reader_thread: object | None = None + self._entered = False + self._closed = False + self._closing = False + self._verdict: str | None = None + self._verdict_lock = threading.Lock() + + def __enter__(self) -> KeyboardStop: + if not self.is_main_thread(): + raise RuntimeError("KeyboardStop must be entered on the main thread") + if self._entered and not self._closed: + raise RuntimeError("KeyboardStop is already active") + if self._closed: + raise RuntimeError("KeyboardStop contexts cannot be reused") + self._entered = True + try: + self._install_signal_handlers() + self._start_terminal_reader() + except Exception: + self._cleanup_entry() + raise + return self + + def __exit__(self, exc_type: object, exc: object, traceback: object) -> None: + if self._closed: + return + self._cleanup_entry() + + def _cleanup_entry(self) -> None: + """Restore every entry side effect; safe after partial setup.""" + + self._closed = True + self._closing = True + reader = self._reader_thread + join = getattr(reader, "join", None) + if callable(join): + try: + join(timeout=0.5) + except Exception as join_error: + self._warn(f"Keyboard-stop reader cleanup warning: {join_error}") + self._restore_terminal() + self._restore_signal_handlers() + + def stop(self) -> None: + """Request a stop; repeated requests are deliberately harmless.""" + + self.event.set() + + def is_set(self) -> bool: + return self.event.is_set() + + def verdict(self) -> str | None: + """Return the current operator verdict without waiting for input.""" + + with self._verdict_lock: + return self._verdict + + def _install_signal_handlers(self) -> None: + for name in ("SIGINT", "SIGTERM"): + signum = getattr(self.signal_api, name, None) + if signum is None: + continue + previous = self.signal_api.getsignal(signum) + self.signal_api.signal(signum, self._handle_signal) + self._prior_handlers[signum] = previous + + def _restore_signal_handlers(self) -> None: + if not self.is_main_thread(): + self._prior_handlers.clear() + return + for signum, previous in tuple(self._prior_handlers.items()): + try: + self.signal_api.signal(signum, previous) + except Exception as exc: + self._warn(f"Could not restore signal handler {signum}: {exc}") + self._prior_handlers.clear() + + def _handle_signal(self, signum: int, frame: object) -> None: + del signum, frame + self.stop() + + def _start_terminal_reader(self) -> None: + try: + is_tty = bool(self.stdin.isatty()) + except Exception: + is_tty = False + if not is_tty: + self._warn( + "WARNING: keyboard stop is unavailable in non-TTY mode; " + "SIGINT/SIGTERM and the physical e-stop remain available." + ) + return + + fd = self.stdin.fileno() + saved_state = self.termios_api.tcgetattr(fd) + self._terminal_fd = fd + self._terminal_state = saved_state + self.tty_api.setcbreak(fd) + self._reader_thread = self.thread_factory( + target=self._read_terminal, + name="policy-rollout-keyboard-stop", + daemon=True, + ) + self._reader_thread.start() + + def _read_terminal(self) -> None: + while not self._closing and not self.event.is_set(): + try: + readable, _writable, _errors = self.select_fn( + [self.stdin], [], [], POLL_INTERVAL_S + ) + if not readable: + continue + character = self.stdin.read(1) + except Exception as exc: + self._warn(f"WARNING: terminal keyboard reader stopped: {exc}") + return + if character == "": + return + if character in STOP_KEYS: + self.stop() + return + if character in SUCCESS_KEYS: + with self._verdict_lock: + if self._verdict is None: + self._verdict = "success" + continue + if character in FAILURE_KEYS: + with self._verdict_lock: + if self._verdict is None: + self._verdict = "failure" + continue + + def _restore_terminal(self) -> None: + if self._terminal_fd is None or self._terminal_state is None: + return + fd = self._terminal_fd + state = self._terminal_state + self._terminal_fd = None + self._terminal_state = None + try: + self.termios_api.tcsetattr(fd, self.termios_api.TCSADRAIN, state) + except Exception as exc: + self._warn(f"Could not restore terminal state: {exc}") + + def _warn(self, message: str) -> None: + try: + print(message, file=self.warning_stream, flush=True) + except Exception: + pass diff --git a/p3_vlm_orchestrator/policy_rollout/lerobot_policy.py b/p3_vlm_orchestrator/policy_rollout/lerobot_policy.py new file mode 100644 index 0000000..e9a8c1f --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/lerobot_policy.py @@ -0,0 +1,406 @@ +"""Lazy, policy-generic adapter for saved LeRobot checkpoints.""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +import json +from pathlib import Path +import platform +from threading import Lock +from typing import Any + +import numpy as np + +from rebot_operator_kit.rollout.checkpoint import CheckpointBundle +from rebot_operator_kit.rollout.contracts import RolloutObservation + + +class LeRobotCompatibilityError(RuntimeError): + """The installed LeRobot runtime cannot read a checkpoint artifact.""" + + +_PLUGIN_REGISTRATION_LOCK = Lock() + + +def _register_policy_plugins(registrar: Callable[[], None]) -> None: + """Run LeRobot's registrar without importing offline-irrelevant hardware.""" + + import importlib.metadata + + with _PLUGIN_REGISTRATION_LOCK: + discover = importlib.metadata.distributions + policy_distributions = tuple( + distribution + for distribution in discover() + if isinstance(distribution.metadata.get("Name"), str) + and distribution.metadata["Name"].replace("-", "_").startswith( + "lerobot_policy_" + ) + ) + + def discover_policies(*args: Any, **kwargs: Any) -> tuple[Any, ...]: + return policy_distributions + + importlib.metadata.distributions = discover_policies + try: + registrar() + finally: + importlib.metadata.distributions = discover + + +@dataclass(frozen=True) +class _LeRobotAPI: + """Small injectable boundary around LeRobot's versioned public APIs.""" + + config_from_pretrained: Callable[..., Any] + get_policy_class: Callable[[str], type] + make_pre_post_processors: Callable[..., tuple[Any, Any]] + prepare_observation: Callable[[dict[str, np.ndarray], str, str], dict[str, Any]] + inference_mode: Callable[[], Any] + runtime_description: str + + +@dataclass(frozen=True) +class _LoadedBackend: + config: Any + policy: Any + preprocessor: Any + postprocessor: Any + prepare_observation: Callable[[dict[str, np.ndarray], str, str], dict[str, Any]] + inference_mode: Callable[[], Any] + + +def _import_lerobot_api() -> _LeRobotAPI: + """Import all optional inference dependencies at the loading boundary.""" + + try: + import importlib.metadata + + import torch + from lerobot.configs.policies import PreTrainedConfig + from lerobot.policies.factory import ( + get_policy_class, + make_pre_post_processors, + ) + from lerobot.policies.utils import prepare_observation_for_inference + from lerobot.utils.import_utils import register_third_party_plugins + + _register_policy_plugins(register_third_party_plugins) + except Exception as exc: + raise LeRobotCompatibilityError( + "LeRobot inference dependencies are unavailable. Install the same " + "LeRobot checkout and policy plugins used for training. " + f"Original error: {type(exc).__name__}: {exc}" + ) from exc + + try: + version = importlib.metadata.version("lerobot") + except importlib.metadata.PackageNotFoundError: + version = "unknown" + runtime = f"LeRobot {version} on Python {platform.python_version()}" + + return _LeRobotAPI( + config_from_pretrained=PreTrainedConfig.from_pretrained, + get_policy_class=get_policy_class, + make_pre_post_processors=make_pre_post_processors, + prepare_observation=prepare_observation_for_inference, + inference_mode=torch.inference_mode, + runtime_description=runtime, + ) + + +def _compatibility_error( + *, stage: str, checkpoint: Path, api: _LeRobotAPI, error: Exception +) -> LeRobotCompatibilityError: + return LeRobotCompatibilityError( + f"Checkpoint compatibility error during {stage}: {api.runtime_description} " + f"cannot read {checkpoint}. Use the same official LeRobot/Python environment " + f"that produced this checkpoint. Original error: " + f"{type(error).__name__}: {error}" + ) + + +def _saved_preprocessor_device_overrides( + checkpoint: Path, device: str +) -> dict[str, dict[str, str]]: + """Retarget saved processor steps that explicitly carry a device setting.""" + + config_path = checkpoint / "preprocessor_config.json" + if not config_path.is_file(): + return {} + try: + document = json.loads(config_path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError): + return {} + if not isinstance(document, dict) or not isinstance(document.get("steps"), list): + return {} + + overrides: dict[str, dict[str, str]] = {} + for entry in document["steps"]: + if not isinstance(entry, dict): + continue + saved_config = entry.get("config") + if not isinstance(saved_config, dict) or "device" not in saved_config: + continue + key = entry.get("registry_name") + if not isinstance(key, str): + class_path = entry.get("class") + if isinstance(class_path, str): + key = class_path.rsplit(".", 1)[-1] + if isinstance(key, str): + overrides[key] = {"device": device} + return overrides + + +def _validate_local_processor_artifacts(checkpoint: Path) -> None: + """Reject incomplete local pipelines before LeRobot can try a Hub fallback.""" + + checkpoint_root = checkpoint.resolve() + for filename in ("preprocessor_config.json", "postprocessor_config.json"): + config_path = checkpoint / filename + if not config_path.is_file(): + continue + try: + document = json.loads(config_path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + raise ValueError(f"saved processor config {filename} is unreadable: {exc}") from exc + if not isinstance(document, dict) or not isinstance(document.get("steps"), list): + raise ValueError(f"saved processor config {filename} has an invalid steps schema") + for index, entry in enumerate(document["steps"]): + if not isinstance(entry, dict): + raise ValueError( + f"saved processor config {filename} step {index} must be an object" + ) + state_file = entry.get("state_file") + if state_file is None: + continue + if not isinstance(state_file, str) or not state_file: + raise ValueError( + f"saved processor config {filename} step {index} has an invalid state_file" + ) + state_path = (checkpoint / state_file).resolve() + if not state_path.is_relative_to(checkpoint_root): + raise ValueError( + f"saved processor state file must stay inside checkpoint: {state_file}" + ) + if not state_path.is_file(): + raise ValueError( + f"saved processor state file is missing: {state_file}" + ) + + +def _reset_if_supported(component: Any) -> None: + reset = getattr(component, "reset", None) + if callable(reset): + reset() + + +def _load_lerobot_backend(checkpoint: Path, device: str) -> _LoadedBackend: + api = _import_lerobot_api() + checkpoint_text = str(checkpoint) + + try: + config = api.config_from_pretrained( + checkpoint_text, + local_files_only=True, + ) + except Exception as exc: + raise _compatibility_error( + stage="policy configuration loading", + checkpoint=checkpoint, + api=api, + error=exc, + ) from exc + + try: + config.device = device + policy_class = api.get_policy_class(config.type) + except Exception as exc: + raise _compatibility_error( + stage=f"registered policy {getattr(config, 'type', None)!r} resolution", + checkpoint=checkpoint, + api=api, + error=exc, + ) from exc + + try: + policy = policy_class.from_pretrained( + checkpoint_text, + config=config, + local_files_only=True, + ) + to_device = getattr(policy, "to", None) + if callable(to_device): + to_device(device) + evaluate = getattr(policy, "eval", None) + if callable(evaluate): + evaluate() + _reset_if_supported(policy) + except Exception as exc: + raise _compatibility_error( + stage="policy weight loading", + checkpoint=checkpoint, + api=api, + error=exc, + ) from exc + + processor_kwargs: dict[str, Any] = { + "policy_cfg": config, + "pretrained_path": checkpoint_text, + "preprocessor_config_filename": "preprocessor_config.json", + "postprocessor_config_filename": "postprocessor_config.json", + } + device_overrides = _saved_preprocessor_device_overrides(checkpoint, device) + if device_overrides: + processor_kwargs["preprocessor_overrides"] = device_overrides + try: + _validate_local_processor_artifacts(checkpoint) + preprocessor, postprocessor = api.make_pre_post_processors( + **processor_kwargs + ) + _reset_if_supported(preprocessor) + _reset_if_supported(postprocessor) + except Exception as exc: + raise _compatibility_error( + stage=( + "saved processor schema loading " + "(preprocessor_config.json and postprocessor_config.json)" + ), + checkpoint=checkpoint, + api=api, + error=exc, + ) from exc + + return _LoadedBackend( + config=config, + policy=policy, + preprocessor=preprocessor, + postprocessor=postprocessor, + prepare_observation=api.prepare_observation, + inference_mode=api.inference_mode, + ) + + +class LeRobotPolicyAdapter: + """Implement the rollout PolicyAdapter protocol for a LeRobot policy.""" + + def __init__( + self, + *, + bundle: CheckpointBundle, + device: str, + backend: _LoadedBackend, + ) -> None: + self.bundle = bundle + self.device = device + self.config = backend.config + self.policy = backend.policy + self.preprocessor = backend.preprocessor + self.postprocessor = backend.postprocessor + self._prepare_observation = backend.prepare_observation + self._inference_mode = backend.inference_mode + + @classmethod + def from_checkpoint( + cls, bundle: CheckpointBundle, device: str + ) -> LeRobotPolicyAdapter: + backend = _load_lerobot_backend(bundle.path, device) + return cls(bundle=bundle, device=device, backend=backend) + + def reset(self) -> None: + """Clear policy and saved-processor episode state before a new trial.""" + + for label, component in ( + ("policy", self.policy), + ("preprocessor", self.preprocessor), + ("postprocessor", self.postprocessor), + ): + try: + _reset_if_supported(component) + except Exception as exc: + raise RuntimeError(f"{label} reset failed: {exc}") from exc + + def predict(self, observation: RolloutObservation) -> np.ndarray: + try: + state = np.array( + observation.state_deg, + dtype=np.float32, + copy=True, + ) + except (TypeError, ValueError) as exc: + raise ValueError("LeRobot observation state must be numeric") from exc + expected_state_shape = (self.bundle.action_dimension,) + if state.shape != expected_state_shape: + raise ValueError( + "LeRobot observation state must have shape " + f"{expected_state_shape}; received {state.shape}" + ) + + raw_observation = { + "observation.images.front": np.array(observation.front, copy=True), + "observation.images.side": np.array(observation.side, copy=True), + "observation.state": state, + } + with self._inference_mode(): + prepared = self._prepare_observation( + raw_observation, + self.device, + self.bundle.task, + ) + prepared.pop("robot_type", None) + if prepared.get("task") != self.bundle.task: + raise ValueError( + "LeRobot prepared task does not match the checkpoint task" + ) + expected_keys = ( + "observation.images.front", + "observation.images.side", + "observation.state", + "task", + ) + if tuple(prepared) != expected_keys: + raise ValueError( + "LeRobot prepared observation keys must be exactly " + f"{expected_keys}; received {tuple(prepared)}" + ) + processed_observation = self.preprocessor(prepared) + action_chunk = self.policy.predict_action_chunk(processed_observation) + postprocessed = self.postprocessor(action_chunk) + + value = postprocessed + for method_name in ("detach", "cpu"): + method = getattr(value, method_name, None) + if callable(method): + value = method() + to_numpy = getattr(value, "numpy", None) + if callable(to_numpy): + value = to_numpy() + try: + array = np.array(value, dtype=np.float64, copy=True) + except (TypeError, ValueError) as exc: + raise ValueError( + "LeRobot postprocessed action chunk must be numeric" + ) from exc + + if array.ndim == 3: + if array.shape[0] != 1: + raise ValueError( + "LeRobot postprocessed action chunk batch dimension must be 1; " + f"received {array.shape}" + ) + array = array[0] + if array.ndim != 2 or array.shape[1] != self.bundle.action_dimension: + raise ValueError( + "LeRobot postprocessed action chunk shape must be [steps, 7]; " + f"received {array.shape}" + ) + if array.shape[0] < 1: + raise ValueError( + "LeRobot postprocessed action chunk must contain at least one step" + ) + if not np.isfinite(array).all(): + raise ValueError( + "LeRobot postprocessed action chunk contains nonfinite values" + ) + return np.array(array, dtype=np.float64, copy=True) diff --git a/p3_vlm_orchestrator/policy_rollout/offline.py b/p3_vlm_orchestrator/policy_rollout/offline.py new file mode 100644 index 0000000..0fde9a7 --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/offline.py @@ -0,0 +1,287 @@ +"""Offline-only checkpoint evaluation over finalized LeRobot datasets.""" + +from __future__ import annotations + +import argparse +from collections.abc import Callable, Sequence +from dataclasses import dataclass +import json +from pathlib import Path +import sys +import time +from typing import Any, TextIO + +import numpy as np + +from p3_vlm_orchestrator.policy_rollout.lerobot_policy import LeRobotPolicyAdapter +from rebot_operator_kit.rollout.checkpoint import CheckpointBundle +from rebot_operator_kit.rollout.contracts import RolloutObservation + + +@dataclass(frozen=True) +class OfflinePrediction: + episode_index: int + sample_index: int + shape: tuple[int, ...] + minimum: float + maximum: float + latency_s: float + + +def _cpu_numpy(value: Any) -> np.ndarray: + converted = value + for method_name in ("detach", "cpu"): + method = getattr(converted, method_name, None) + if callable(method): + converted = method() + to_numpy = getattr(converted, "numpy", None) + if callable(to_numpy): + converted = to_numpy() + return np.array(converted, copy=True) + + +def _scalar(value: Any, *, label: str) -> float: + array = _cpu_numpy(value) + if array.size != 1: + raise ValueError(f"dataset {label} must be scalar; received {array.shape}") + try: + return float(array.reshape(()).item()) + except (TypeError, ValueError) as exc: + raise ValueError(f"dataset {label} must be numeric") from exc + + +def _camera_array(value: Any, *, key: str) -> np.ndarray: + array = _cpu_numpy(value) + if array.ndim != 3: + raise ValueError(f"dataset {key} must be a three-dimensional image") + + channels = (1, 3, 4) + if array.shape[-1] in channels and array.shape[0] not in channels: + image = array + elif array.shape[0] in channels: + image = np.moveaxis(array, 0, -1) + elif array.shape[-1] in channels: + image = array + else: + raise ValueError(f"dataset {key} must be CHW or HWC") + + if np.issubdtype(image.dtype, np.floating): + if not np.isfinite(image).all(): + raise ValueError(f"dataset {key} contains nonfinite pixels") + minimum = float(image.min()) + maximum = float(image.max()) + if minimum >= 0.0 and maximum <= 1.0: + image = np.rint(image * 255.0) + elif minimum >= 0.0 and maximum <= 255.0: + image = np.rint(image) + else: + raise ValueError(f"dataset {key} pixels must be in [0, 1] or [0, 255]") + elif np.issubdtype(image.dtype, np.integer): + if image.size and (int(image.min()) < 0 or int(image.max()) > 255): + raise ValueError(f"dataset {key} pixels must be in [0, 255]") + else: + raise ValueError(f"dataset {key} must contain numeric pixels") + return np.array(image, dtype=np.uint8, copy=True) + + +def _state_array(value: Any) -> np.ndarray: + try: + state = np.array(_cpu_numpy(value), dtype=np.float32, copy=True) + except (TypeError, ValueError) as exc: + raise ValueError("dataset observation.state must be numeric") from exc + if state.shape != (7,): + raise ValueError( + "dataset observation.state must have shape (7,); " + f"received {state.shape}" + ) + if not np.isfinite(state).all(): + raise ValueError("dataset observation.state contains nonfinite values") + return state + + +def _load_lerobot_dataset(root: Path) -> Any: + """Load only a complete local dataset; LeRobot remains an optional import.""" + + info_path = root / "meta" / "info.json" + if not root.is_dir() or not info_path.is_file(): + raise ValueError(f"finalized LeRobot dataset root is missing: {root}") + required_paths = ( + root / "meta" / "stats.json", + root / "meta" / "tasks.parquet", + root / "meta" / "episodes", + root / "data", + ) + missing = [str(path.relative_to(root)) for path in required_paths if not path.exists()] + if missing: + raise ValueError( + "finalized LeRobot dataset is incomplete; missing " + ", ".join(missing) + ) + try: + info = json.loads(info_path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + raise ValueError(f"finalized LeRobot dataset info is unreadable: {exc}") from exc + if not isinstance(info, dict): + raise ValueError("finalized LeRobot dataset info must be an object") + if not isinstance(info.get("total_episodes"), int) or info["total_episodes"] < 1: + raise ValueError("finalized LeRobot dataset contains no episodes") + if not isinstance(info.get("total_frames"), int) or info["total_frames"] < 1: + raise ValueError("finalized LeRobot dataset contains no frames") + + try: + from lerobot.datasets.lerobot_dataset import LeRobotDataset + except Exception as exc: + raise RuntimeError( + "Offline evaluation requires the same LeRobot environment used to " + f"record the dataset. Original error: {type(exc).__name__}: {exc}" + ) from exc + + repo_id = info.get("repo_id") + if not isinstance(repo_id, str) or not repo_id.strip(): + repo_id = f"local/{root.name}" + return LeRobotDataset(repo_id=repo_id, root=root) + + +def evaluate_checkpoint( + bundle: CheckpointBundle, + dataset_root: Path, + *, + episodes: int, + device: str = "cpu", + dataset_loader: Callable[[Path], Any] | None = None, + adapter_factory: Callable[[CheckpointBundle, str], Any] | None = None, + clock: Callable[[], float] = time.perf_counter, + output: TextIO | None = None, +) -> list[OfflinePrediction]: + """Evaluate every sample belonging to exactly N distinct episodes.""" + + if isinstance(episodes, bool) or not isinstance(episodes, int) or episodes <= 0: + raise ValueError("episodes must be a positive integer") + root = Path(dataset_root).expanduser().resolve() + load_dataset = dataset_loader or _load_lerobot_dataset + make_adapter = adapter_factory or LeRobotPolicyAdapter.from_checkpoint + stream = output or sys.stdout + + dataset = load_dataset(root) + adapter = make_adapter(bundle, device) + selected_episodes: list[int] = [] + selected_set: set[int] = set() + active_episode: int | None = None + results: list[OfflinePrediction] = [] + + for sample_index in range(len(dataset)): + sample = dataset[sample_index] + if not isinstance(sample, dict): + raise ValueError(f"dataset sample {sample_index} must be a mapping") + if "episode_index" not in sample: + raise ValueError(f"dataset sample {sample_index} has no episode_index") + episode_index = int( + _scalar(sample["episode_index"], label="episode_index") + ) + if episode_index not in selected_set: + if len(selected_episodes) >= episodes: + continue + selected_episodes.append(episode_index) + selected_set.add(episode_index) + + if episode_index != active_episode: + reset = getattr(adapter, "reset", None) + if callable(reset): + reset() + active_episode = episode_index + + task = sample.get("task") + if task != bundle.task: + raise ValueError( + "dataset task mismatch at " + f"episode {episode_index}, sample {sample_index}: " + f"expected {bundle.task!r}, received {task!r}" + ) + for key in ( + "observation.images.front", + "observation.images.side", + "observation.state", + ): + if key not in sample: + raise ValueError(f"dataset sample {sample_index} is missing {key}") + + captured_s = ( + _scalar(sample["timestamp"], label="timestamp") + if "timestamp" in sample + else 0.0 + ) + observation = RolloutObservation( + front=_camera_array( + sample["observation.images.front"], + key="observation.images.front", + ), + side=_camera_array( + sample["observation.images.side"], + key="observation.images.side", + ), + state_deg=_state_array(sample["observation.state"]), + task=bundle.task, + captured_monotonic_s=captured_s, + ) + + started = clock() + prediction = np.asarray(adapter.predict(observation), dtype=np.float64) + finished = clock() + if prediction.ndim != 2 or prediction.shape[0] < 1 or prediction.shape[1] != 7: + raise ValueError( + "offline policy prediction must have shape [steps, 7]; " + f"received {prediction.shape}" + ) + if not np.isfinite(prediction).all(): + raise ValueError("offline policy prediction contains nonfinite values") + result = OfflinePrediction( + episode_index=episode_index, + sample_index=sample_index, + shape=tuple(prediction.shape), + minimum=float(prediction.min()), + maximum=float(prediction.max()), + latency_s=max(0.0, finished - started), + ) + results.append(result) + print( + f"episode={result.episode_index} sample={result.sample_index} " + f"shape={result.shape} min={result.minimum:.6f} " + f"max={result.maximum:.6f} latency_ms={result.latency_s * 1000:.3f}", + file=stream, + ) + + if len(selected_episodes) != episodes: + raise ValueError( + f"requested {episodes} distinct episodes but found " + f"{len(selected_episodes)} in the dataset" + ) + + return results + + +def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Run a saved LeRobot checkpoint over held-out recorded episodes." + ) + parser.add_argument("--checkpoint", required=True, type=Path) + parser.add_argument("--dataset", required=True, type=Path) + parser.add_argument("--episodes", type=int, default=2) + parser.add_argument("--device", default="cpu") + return parser.parse_args(argv) + + +def main(argv: Sequence[str] | None = None) -> int: + args = _parse_args(argv) + bundle = CheckpointBundle.load(args.checkpoint) + results = evaluate_checkpoint( + bundle, + args.dataset, + episodes=args.episodes, + device=args.device, + ) + if not results: + raise RuntimeError("No samples were found in the selected dataset episodes") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/p3_vlm_orchestrator/policy_rollout/rebot_robot.py b/p3_vlm_orchestrator/policy_rollout/rebot_robot.py new file mode 100644 index 0000000..98ff366 --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/rebot_robot.py @@ -0,0 +1,877 @@ +"""Authenticated LeRobot hardware adapter for policy rollout. + +All serial, camera, LeRobot, and motor-driver imports are behind the default +backend seam. Importing this module is therefore safe on machines without the +physical runtime installed. +""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable, Mapping, Sequence +from dataclasses import fields, is_dataclass +from hashlib import sha256 +import math +from numbers import Real +from pathlib import Path +import sys +import time +from typing import Any, Protocol + +import numpy as np + +from rebot_operator_kit.rollout.contracts import RolloutObservation + + +JOINT_COUNT = 7 +FIRST_LIVE_SPEED_SCALE_MIN = 0.10 +FIRST_LIVE_SPEED_SCALE_MAX = 0.20 +MAX_RELATIVE_TARGET_DEG = 1.5 +ACTION_MATCH_TOLERANCE_DEG = 1e-9 +FOLLOWER_RUNTIME_CONTRACTS = ( + "follower_driver_contract", + "follower_base_implementation", + "follower_dm_implementation", +) +EXPECTED_CAMERA_KEYS = ("front", "side") +ALLOWED_CAMERA_METADATA_KEYS = ( + "excluded_screen_index", + "excluded_screen_name", + "minimum_measured_fps", +) +EXPECTED_RECORDING_KEYS = ( + "observation.images.front", + "observation.images.side", +) +CALIBRATION_FIELDS = ( + "id", + "drive_mode", + "homing_offset", + "range_min", + "range_max", +) + + +class FollowerCleanupError(RuntimeError): + """Aggregates best-effort follower resource cleanup failures.""" + + def __init__(self, errors: Sequence[Exception]) -> None: + self.errors = tuple(errors) + super().__init__( + f"Follower cleanup encountered {len(self.errors)} resource error(s)" + ) + + +class _Backend(Protocol): + def make_camera_config(self, **kwargs: object) -> object: ... + + def make_follower_config(self, **kwargs: object) -> object: ... + + def make_follower(self, config: object) -> object: ... + + +class ReBotPolicyRobot: + """RobotAdapter over the verified seven-joint LeRobot follower plugin.""" + + def __init__( + self, + *, + follower: object, + joint_names: Sequence[str], + task: str, + camera_contract: Mapping[str, Mapping[str, object]], + calibration_path: Path, + monotonic_clock: Callable[[], float] = time.monotonic, + ) -> None: + """Bind an already verified plugin; ``from_profile`` performs verification.""" + + self.follower = follower + self.joint_names = tuple(joint_names) + self.feature_names = tuple(f"{name}.pos" for name in self.joint_names) + self.task = task + self.camera_contract = { + key: dict(value) for key, value in camera_contract.items() + } + self.calibration_path = calibration_path + self.monotonic_clock = monotonic_clock + self._connected = False + + @classmethod + def from_profile( + cls, + *, + profile_snapshot: Mapping[str, Any], + runtime_root: Path | str, + speed_scale: float, + monotonic_clock: Callable[[], float] = time.monotonic, + backend: _Backend | None = None, + serial_ports_provider: Callable[[], Iterable[object]] | None = None, + ) -> ReBotPolicyRobot: + """Validate the authenticated profile and build without connecting.""" + + checked_speed = _finite_float(speed_scale, "speed_scale") + if not ( + FIRST_LIVE_SPEED_SCALE_MIN + <= checked_speed + <= FIRST_LIVE_SPEED_SCALE_MAX + ): + raise ValueError( + "speed_scale must be within [0.10, 0.20] for a live-capable follower" + ) + if not isinstance(profile_snapshot, Mapping): + raise ValueError("Authenticated profile snapshot must be an object") + + root = Path(runtime_root).expanduser().resolve() + calibration_section = _required_mapping( + profile_snapshot, "calibration", "profile" + ) + follower_profile = _required_mapping( + calibration_section, "follower", "profile calibration" + ) + follower_type = _required_text( + follower_profile.get("type"), "Follower calibration type" + ) + follower_id = _required_text( + follower_profile.get("id"), "Follower calibration id" + ) + calibration_path = _verified_runtime_file( + root, follower_profile, "Follower calibration" + ) + for contract_name in FOLLOWER_RUNTIME_CONTRACTS: + entry = _required_mapping( + calibration_section, contract_name, "profile calibration" + ) + _verified_runtime_file(root, entry, contract_name) + + coordinates = _required_mapping( + profile_snapshot, "coordinate_contract", "profile" + ) + joints = _profile_joints(coordinates) + joint_names = tuple(joint["name"] for joint in joints) + expected_directions = { + joint["name"]: joint["direction"] for joint in joints + } + expected_limits = {joint["name"]: joint["limits"] for joint in joints} + expected_features = tuple(joint["feature"] for joint in joints) + if expected_features != tuple(f"{name}.pos" for name in joint_names): + raise ValueError("Profile joint features must match the seven physical joints") + + calibration_values = _load_calibration(calibration_path, joint_names) + camera_contract = _camera_contract(profile_snapshot) + task = _required_text( + _required_mapping( + profile_snapshot, "collection_defaults", "profile" + ).get("task"), + "Authenticated rollout task", + ) + hardware = _required_mapping( + profile_snapshot, "hardware_identity", "profile" + ) + follower_usb = _required_mapping( + hardware, "follower_usb", "profile hardware identity" + ) + expected_vid = _usb_id(follower_usb.get("vid"), "follower USB VID") + expected_pid = _usb_id(follower_usb.get("pid"), "follower USB PID") + + if serial_ports_provider is None: + serial_ports_provider = _default_serial_ports + try: + ports = list(serial_ports_provider()) + except Exception as exc: + raise ValueError(f"Follower USB discovery failed: {exc}") from exc + matches = [ + port + for port in ports + if getattr(port, "vid", None) == expected_vid + and getattr(port, "pid", None) == expected_pid + ] + if len(matches) != 1: + devices = [str(getattr(port, "device", "")) for port in matches] + raise ValueError( + "Follower USB discovery expected exactly one authenticated device; " + f"found {devices}" + ) + device = getattr(matches[0], "device", None) + if not isinstance(device, str) or not device: + raise ValueError("Authenticated follower USB device has no usable port") + + checked_backend: _Backend = backend or _DefaultBackend(root) + cameras = { + key: checked_backend.make_camera_config( + index_or_path=contract["index"], + fps=contract["fps"], + width=contract["width"], + height=contract["height"], + fourcc="MJPG", + ) + for key, contract in camera_contract.items() + } + collection = _required_mapping( + profile_snapshot, "collection_defaults", "profile" + ) + profile_velocity = collection.get("motor_velocity") + if ( + isinstance(profile_velocity, bool) + or not isinstance(profile_velocity, (int, float)) + or not math.isfinite(float(profile_velocity)) + or float(profile_velocity) != 2000.0 + ): + raise ValueError( + "Profile motor velocity must equal exactly 2000.0 degrees/s" + ) + gripper_force = _finite_float( + collection.get("gripper_force", 0.05), + "Profile gripper force", + ) + if not 0.0 <= gripper_force <= 1.0: + raise ValueError("Profile gripper force must be in [0, 1]") + + expected_velocity = [2000.0 * checked_speed] * JOINT_COUNT + follower_config = checked_backend.make_follower_config( + port=device, + id=follower_id, + calibration_dir=calibration_path.parent, + can_adapter="damiao", + dm_serial_baud=921600, + max_relative_target=MAX_RELATIVE_TARGET_DEG, + pos_vel_velocity=expected_velocity, + force_pos_torque_ration=gripper_force, + disable_torque_on_disconnect=True, + cameras=cameras, + ) + follower = checked_backend.make_follower(follower_config) + _verify_plugin_binding( + follower=follower, + follower_type=follower_type, + joint_names=joint_names, + expected_directions=expected_directions, + expected_limits=expected_limits, + camera_contract=camera_contract, + calibration_path=calibration_path, + calibration_values=calibration_values, + expected_velocity=expected_velocity, + ) + + return cls( + follower=follower, + joint_names=joint_names, + task=task, + camera_contract=camera_contract, + calibration_path=calibration_path, + monotonic_clock=monotonic_clock, + ) + + def connect(self) -> None: + if self._connected: + raise RuntimeError("ReBot policy follower is already connected") + try: + self.follower.connect(calibrate=False) + if not bool(getattr(self.follower, "is_connected", False)): + raise RuntimeError("Follower plugin did not report a connected state") + except Exception as connect_error: + cleanup_error = _cleanup_follower_resources(self.follower) + self._connected = False + if cleanup_error is not None: + raise connect_error from cleanup_error + raise + self._connected = True + + def disconnect(self) -> None: + aggregate_connected = bool(getattr(self.follower, "is_connected", False)) + if aggregate_connected: + try: + self.follower.disconnect() + finally: + self._connected = False + return + if not self._connected and not _follower_has_resources(self.follower): + return + + cleanup_error = _cleanup_follower_resources(self.follower) + self._connected = False + if cleanup_error is not None: + raise cleanup_error from cleanup_error.errors[0] + + def observe(self) -> RolloutObservation: + self._require_connected() + try: + raw = self.follower.get_observation() + except Exception: + raise + if not isinstance(raw, Mapping): + raise ValueError("Follower observation must be a mapping") + + images: dict[str, np.ndarray] = {} + for key in EXPECTED_CAMERA_KEYS: + if key not in raw: + raise ValueError(f"Follower observation is missing camera {key!r}") + try: + image = np.asarray(raw[key]) + except Exception as exc: + raise ValueError(f"Follower observation camera {key!r} is malformed") from exc + contract = self.camera_contract[key] + expected_shape = ( + int(contract["height"]), + int(contract["width"]), + 3, + ) + if image.shape != expected_shape or not np.issubdtype( + image.dtype, np.number + ): + raise ValueError( + f"Follower observation camera {key!r} does not match {expected_shape}" + ) + images[key] = image.copy() + + try: + state = np.asarray([raw[name] for name in self.feature_names], dtype=float) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError( + "Follower observation is missing or has malformed physical joint positions" + ) from exc + if state.shape != (JOINT_COUNT,) or not np.all(np.isfinite(state)): + raise ValueError( + "Follower observation physical joint positions must be seven finite values" + ) + captured_monotonic_s = _finite_float( + self.monotonic_clock(), + "Follower observation capture monotonic time", + ) + return RolloutObservation( + front=images["front"], + side=images["side"], + state_deg=state.copy(), + task=self.task, + captured_monotonic_s=captured_monotonic_s, + ) + + def send_action(self, action_deg: np.ndarray) -> np.ndarray: + """Invert plugin scales so its returned action is requested physical space.""" + + self._require_connected() + try: + requested = np.asarray(action_deg, dtype=float).copy() + except (TypeError, ValueError) as exc: + raise ValueError("Physical follower action must be seven finite values") from exc + if requested.shape != (JOINT_COUNT,) or not np.all(np.isfinite(requested)): + raise ValueError("Physical follower action must be seven finite values") + + directions = getattr(self.follower.config, "joint_directions", None) + limits = getattr(self.follower.config, "joint_limits", None) + if not isinstance(directions, Mapping): + raise ValueError("Follower plugin direction mapping is missing") + if not isinstance(limits, Mapping): + raise ValueError("Follower plugin joint limits are missing") + command: dict[str, float] = {} + for index, (joint_name, feature_name) in enumerate( + zip(self.joint_names, self.feature_names, strict=True) + ): + direction = _finite_float( + directions.get(joint_name), + f"Follower plugin direction for {joint_name}", + ) + if direction == 0.0: + raise ValueError( + f"Follower plugin direction for {joint_name} must be nonzero" + ) + joint_limit = limits.get(joint_name) + try: + low, high = (float(value) for value in joint_limit) + except (TypeError, ValueError) as exc: + raise ValueError( + f"Follower plugin limit for {joint_name} is malformed" + ) from exc + if not math.isfinite(low) or not math.isfinite(high) or low >= high: + raise ValueError( + f"Follower plugin limit for {joint_name} is invalid" + ) + if requested[index] < low or requested[index] > high: + raise ValueError( + f"Physical follower action for {joint_name} is outside plugin limits" + ) + driver_value = float(requested[index]) / direction + if not math.isfinite(driver_value): + raise ValueError( + f"Physical follower action for {joint_name} cannot be safely inverted" + ) + command[feature_name] = driver_value + + returned_raw = self.follower.send_action(command) + if not isinstance(returned_raw, Mapping) or set(returned_raw) != set( + self.feature_names + ): + raise ValueError( + "Follower plugin returned action must contain exactly seven joint keys" + ) + try: + returned = np.asarray( + [returned_raw[name] for name in self.feature_names], dtype=float + ) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError("Follower plugin returned action is malformed") from exc + if returned.shape != (JOINT_COUNT,) or not np.all(np.isfinite(returned)): + raise ValueError("Follower plugin returned action must be seven finite values") + if not np.allclose( + returned, + requested, + rtol=0.0, + atol=ACTION_MATCH_TOLERANCE_DEG, + ): + raise ValueError( + "Follower plugin returned action disagrees with requested physical action" + ) + return returned.copy() + + def _require_connected(self) -> None: + if not self._connected or not bool( + getattr(self.follower, "is_connected", False) + ): + raise RuntimeError("ReBot policy follower is disconnected") + + +class _DefaultBackend: + """Lazy import seam for the real LeRobot follower and OpenCV camera config.""" + + def __init__(self, runtime_root: Path) -> None: + self.runtime_root = runtime_root + + def _ensure_runtime_paths(self) -> None: + for path in ( + self.runtime_root / "lerobot-robot-seeed-b601", + self.runtime_root / "lerobot" / "src", + ): + resolved = str(path.resolve()) + if path.is_dir() and resolved not in sys.path: + sys.path.insert(0, resolved) + + def make_camera_config(self, **kwargs: object) -> object: + self._ensure_runtime_paths() + from lerobot.cameras.opencv.configuration_opencv import OpenCVCameraConfig + + return OpenCVCameraConfig(**kwargs) + + def make_follower_config(self, **kwargs: object) -> object: + self._ensure_runtime_paths() + from lerobot_robot_seeed_b601 import SeeedB601DMFollowerConfig + + return SeeedB601DMFollowerConfig(**kwargs) + + def make_follower(self, config: object) -> object: + self._ensure_runtime_paths() + from lerobot_robot_seeed_b601 import SeeedB601DMFollower + + return SeeedB601DMFollower(config) + + +def _default_serial_ports() -> Iterable[object]: + from serial.tools import list_ports + + return list_ports.comports() + + +def _follower_has_resources(follower: object) -> bool: + if getattr(follower, "bus", None) is not None: + return True + motors = getattr(follower, "motors", None) + if isinstance(motors, Mapping) and bool(motors): + return True + cameras = getattr(follower, "cameras", None) + if isinstance(cameras, Mapping): + for camera in cameras.values(): + try: + if bool(getattr(camera, "is_connected", False)): + return True + except Exception: + return True + return False + + +def _cleanup_follower_resources( + follower: object, +) -> FollowerCleanupError | None: + """Best-effort cleanup that does not depend on aggregate connection state.""" + + errors: list[Exception] = [] + bus = getattr(follower, "bus", None) + if bus is not None: + disable_all = getattr(bus, "disable_all", None) + if callable(disable_all): + try: + disable_all() + except Exception as exc: + errors.append(exc) + + motors = getattr(follower, "motors", None) + if isinstance(motors, Mapping): + for motor in motors.values(): + disable = getattr(motor, "disable", None) + if callable(disable): + try: + disable() + except Exception as exc: + errors.append(exc) + close = getattr(motor, "close", None) + if callable(close): + try: + close() + except Exception as exc: + errors.append(exc) + try: + setattr(follower, "motors", {}) + except Exception as exc: + errors.append(exc) + + if bus is not None: + close_bus = getattr(bus, "close", None) + if callable(close_bus): + try: + close_bus() + except Exception as exc: + errors.append(exc) + try: + setattr(follower, "bus", None) + except Exception as exc: + errors.append(exc) + + cameras = getattr(follower, "cameras", None) + if isinstance(cameras, Mapping): + for camera in cameras.values(): + try: + connected = bool(getattr(camera, "is_connected", False)) + except Exception as exc: + errors.append(exc) + connected = True + if connected: + disconnect = getattr(camera, "disconnect", None) + if callable(disconnect): + try: + disconnect() + except Exception as exc: + errors.append(exc) + + return FollowerCleanupError(errors) if errors else None + + +def _required_mapping( + parent: Mapping[str, Any], key: str, label: str +) -> Mapping[str, Any]: + value = parent.get(key) + if not isinstance(value, Mapping): + raise ValueError(f"Authenticated {label} {key} must be an object") + return value + + +def _required_text(value: object, label: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{label} must be nonempty") + return value + + +def _finite_float(value: object, label: str) -> float: + try: + result = float(value) + except (TypeError, ValueError) as exc: + raise ValueError(f"{label} must be finite") from exc + if not math.isfinite(result): + raise ValueError(f"{label} must be finite") + return result + + +def _usb_id(value: object, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 0xFFFF: + raise ValueError(f"Authenticated {label} is invalid") + return value + + +def _verified_runtime_file( + runtime_root: Path, + entry: Mapping[str, Any], + label: str, +) -> Path: + relative = entry.get("runtime_relative_path") + expected_digest = entry.get("sha256") + if not isinstance(relative, str) or not relative or Path(relative).is_absolute(): + raise ValueError(f"{label} runtime path is invalid") + path = (runtime_root / relative).resolve() + try: + path.relative_to(runtime_root) + except ValueError as exc: + raise ValueError(f"{label} runtime path escapes the authenticated root") from exc + if not path.is_file(): + raise ValueError(f"{label} fingerprint cannot be verified; file is missing") + actual_digest = sha256(path.read_bytes()).hexdigest() + if not isinstance(expected_digest, str) or actual_digest != expected_digest: + raise ValueError(f"{label} fingerprint does not match the authenticated profile") + return path + + +def _profile_joints(coordinates: Mapping[str, Any]) -> list[dict[str, Any]]: + if coordinates.get("action_dimension") != JOINT_COUNT: + raise ValueError("Profile action dimension must be seven") + raw_joints = coordinates.get("joints") + if not isinstance(raw_joints, list) or len(raw_joints) != JOINT_COUNT: + raise ValueError("Profile must define exactly seven physical follower joints") + joints: list[dict[str, Any]] = [] + for raw in raw_joints: + if not isinstance(raw, Mapping): + raise ValueError("Profile physical follower joint entry is malformed") + name = _required_text(raw.get("name"), "Profile joint name") + feature = _required_text(raw.get("feature"), f"Profile feature for {name}") + direction = _finite_float( + raw.get("leader_to_follower_scale"), + f"Profile direction for {name}", + ) + if direction == 0.0: + raise ValueError(f"Profile direction for {name} must be nonzero") + raw_limits = raw.get("soft_limit_degrees") + try: + low, high = (float(value) for value in raw_limits) + except (TypeError, ValueError) as exc: + raise ValueError(f"Profile limits for {name} are malformed") from exc + if not math.isfinite(low) or not math.isfinite(high) or low >= high: + raise ValueError(f"Profile limits for {name} are invalid") + joints.append( + { + "name": name, + "feature": feature, + "direction": direction, + "limits": (low, high), + } + ) + names = [joint["name"] for joint in joints] + if len(set(names)) != JOINT_COUNT: + raise ValueError("Profile physical follower joint names must be unique") + return joints + + +def _load_calibration( + path: Path, joint_names: Sequence[str] +) -> dict[str, dict[str, float]]: + import json + + try: + value = json.loads(path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + raise ValueError("Follower calibration is not valid JSON") from exc + return _normalize_calibration( + value, + joint_names, + label="Follower calibration file", + ) + + +def _normalize_calibration( + calibration: object, + joint_names: Sequence[str], + *, + label: str, +) -> dict[str, dict[str, float]]: + if not isinstance(calibration, Mapping) or list(calibration) != list(joint_names): + raise ValueError( + f"{label} joint order does not match the authenticated profile" + ) + return { + joint_name: _normalize_calibration_entry( + calibration[joint_name], + label=f"{label} entry {joint_name}", + ) + for joint_name in joint_names + } + + +def _normalize_calibration_entry( + entry: object, + *, + label: str, +) -> dict[str, float]: + if isinstance(entry, Mapping): + values = dict(entry) + elif hasattr(entry, "__dict__"): + values = { + name: value + for name, value in vars(entry).items() + if not name.startswith("_") + } + elif is_dataclass(entry): + values = {field.name: getattr(entry, field.name) for field in fields(entry)} + else: + slots = getattr(type(entry), "__slots__", ()) + if isinstance(slots, str): + slots = (slots,) + values = { + name: getattr(entry, name) + for name in slots + if isinstance(name, str) and not name.startswith("_") + } + + actual_fields = set(values) + expected_fields = set(CALIBRATION_FIELDS) + if actual_fields != expected_fields: + missing = sorted(expected_fields - actual_fields) + extra = sorted(actual_fields - expected_fields) + raise ValueError( + f"{label} calibration fields mismatch; missing={missing}, extra={extra}" + ) + + normalized: dict[str, float] = {} + for field_name in CALIBRATION_FIELDS: + value = values[field_name] + if isinstance(value, bool) or not isinstance(value, Real): + raise ValueError( + f"{label} calibration field {field_name} must be a finite number" + ) + numeric = float(value) + if not math.isfinite(numeric): + raise ValueError( + f"{label} calibration field {field_name} must be a finite number" + ) + normalized[field_name] = numeric + return normalized + + +def _camera_contract( + profile: Mapping[str, Any], +) -> dict[str, dict[str, int | str]]: + cameras = _required_mapping(profile, "camera_defaults", "profile") + extra_entries = set(cameras) - set(EXPECTED_CAMERA_KEYS) - set( + ALLOWED_CAMERA_METADATA_KEYS + ) + if extra_entries: + raise ValueError( + "Profile camera contract must contain exactly front and side; " + f"extra entries: {sorted(extra_entries)}" + ) + camera_entries = { + key for key, value in cameras.items() if isinstance(value, Mapping) + } + if camera_entries != set(EXPECTED_CAMERA_KEYS): + raise ValueError("Profile camera contract must contain exactly front and side") + result: dict[str, dict[str, int | str]] = {} + for key, expected_recording_key in zip( + EXPECTED_CAMERA_KEYS, EXPECTED_RECORDING_KEYS, strict=True + ): + raw = cameras.get(key) + if not isinstance(raw, Mapping): + raise ValueError(f"Profile camera {key} must be an object") + if raw.get("recording_key") != expected_recording_key: + raise ValueError(f"Profile camera {key} recording key is invalid") + index = raw.get("index") + if isinstance(index, bool) or not isinstance(index, int) or index < 0: + raise ValueError(f"Profile camera {key} index is invalid") + result[key] = { + "recording_key": expected_recording_key, + "index": index, + "width": _positive_integer(raw.get("width"), f"{key} camera width"), + "height": _positive_integer(raw.get("height"), f"{key} camera height"), + "fps": _positive_integer(raw.get("fps"), f"{key} camera FPS"), + } + if result["front"]["index"] == result["side"]["index"]: + raise ValueError("Profile front and side camera indices must be distinct") + return result + + +def _verify_plugin_binding( + *, + follower: object, + follower_type: str, + joint_names: Sequence[str], + expected_directions: Mapping[str, float], + expected_limits: Mapping[str, tuple[float, float]], + camera_contract: Mapping[str, Mapping[str, object]], + calibration_path: Path, + calibration_values: Mapping[str, object], + expected_velocity: Sequence[float], +) -> None: + if getattr(follower, "name", None) != follower_type: + raise ValueError("Instantiated follower plugin type does not match the profile") + if list(getattr(follower, "motor_names", ())) != list(joint_names): + raise ValueError("Instantiated follower plugin joint order does not match the profile") + + plugin_config = getattr(follower, "config", None) + actual_directions_raw = getattr(plugin_config, "joint_directions", None) + actual_limits_raw = getattr(plugin_config, "joint_limits", None) + if not isinstance(actual_directions_raw, Mapping) or not isinstance( + actual_limits_raw, Mapping + ): + raise ValueError("Instantiated follower plugin direction or limit contract is missing") + actual_directions = { + name: _finite_float(value, f"Follower plugin direction for {name}") + for name, value in actual_directions_raw.items() + } + try: + actual_limits = { + name: tuple(float(value) for value in values) + for name, values in actual_limits_raw.items() + } + except (TypeError, ValueError) as exc: + raise ValueError("Instantiated follower plugin limits are malformed") from exc + if actual_directions != dict(expected_directions): + raise ValueError("Instantiated follower plugin directions do not match the profile") + if actual_limits != dict(expected_limits): + raise ValueError("Instantiated follower plugin limits do not match the profile") + if getattr(plugin_config, "max_relative_target", None) != MAX_RELATIVE_TARGET_DEG: + raise ValueError( + "Instantiated follower plugin relative target does not match 1.5 degrees" + ) + if getattr(plugin_config, "disable_torque_on_disconnect", None) is not True: + raise ValueError( + "Instantiated follower plugin torque-off disconnect flag must be exactly True" + ) + try: + actual_velocity = np.asarray( + getattr(plugin_config, "pos_vel_velocity"), dtype=float + ) + except (TypeError, ValueError) as exc: + raise ValueError("Instantiated follower plugin velocity is malformed") from exc + if ( + actual_velocity.shape != (JOINT_COUNT,) + or not np.all(np.isfinite(actual_velocity)) + or not np.array_equal(actual_velocity, np.asarray(expected_velocity, dtype=float)) + ): + raise ValueError("Instantiated follower plugin velocity does not match the profile") + + actual_calibration = Path( + getattr(follower, "calibration_fpath", "") + ).expanduser().resolve() + if actual_calibration != calibration_path: + raise ValueError("Follower plugin calibration path does not match the profile") + loaded_calibration = _normalize_calibration( + getattr(follower, "calibration", None), + joint_names, + label="Follower plugin calibration", + ) + if loaded_calibration != calibration_values: + raise ValueError( + "Follower plugin calibration values do not match the verified calibration file" + ) + + follower_cameras = getattr(follower, "cameras", None) + config_cameras = getattr(plugin_config, "cameras", None) + if not isinstance(follower_cameras, Mapping) or list(follower_cameras) != list( + EXPECTED_CAMERA_KEYS + ): + raise ValueError("Instantiated follower plugin camera keys do not match the profile") + if not isinstance(config_cameras, Mapping) or list(config_cameras) != list( + EXPECTED_CAMERA_KEYS + ): + raise ValueError("Follower plugin camera config does not match the profile") + for key in EXPECTED_CAMERA_KEYS: + actual = config_cameras[key] + expected = camera_contract[key] + for field, actual_field in ( + ("index", "index_or_path"), + ("width", "width"), + ("height", "height"), + ("fps", "fps"), + ): + if getattr(actual, actual_field, None) != expected[field]: + raise ValueError( + f"Follower plugin camera {key} {field} does not match the profile" + ) + if bool(getattr(follower, "is_connected", False)): + raise ValueError("Follower plugin connected during construction validation") + + +def _positive_integer(value: object, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise ValueError(f"{label} must be a positive integer") + return value diff --git a/p3_vlm_orchestrator/policy_rollout/runner.py b/p3_vlm_orchestrator/policy_rollout/runner.py new file mode 100644 index 0000000..210232a --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/runner.py @@ -0,0 +1,757 @@ +"""One-step receding-horizon policy rollout independent of concrete hardware.""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from datetime import datetime, timezone +import json +from pathlib import Path +from typing import IO, Protocol + +import numpy as np + +from rebot_operator_kit.rollout.contracts import ( + PolicyAdapter, + RobotAdapter, + RolloutObservation, +) +from rebot_operator_kit.rollout.safety import SafetyDecision, SafetyGovernor + + +EXPECTED_ACTION_DIMENSION = 7 +EPISODE_TIMEOUT_S = 30.0 +EPISODE_MAX_ACTIONS = 300 +_OPERATOR_VERDICTS = frozenset(("success", "failure")) + + +class _StopEvent(Protocol): + def is_set(self) -> bool: ... + + +class ActionGuard(Protocol): + """Fail-closed geometric validator for one physical follower action.""" + + def validate(self, action_deg: np.ndarray) -> None: ... + + +@dataclass(frozen=True) +class RunSummary: + """Terminal outcome of one rollout run.""" + + mode: str + cycles_completed: int + actions_attempted: int + actions_confirmed: int + terminal_reason: str + primary_fault_reason: str | None + cleanup_fault_reason: str | None + audit_fault_reason: str | None + attempt: int = 1 + elapsed_seconds: float = 0.0 + clamp_count: int = 0 + + @property + def actions_sent(self) -> int: + """Backward-compatible alias for confirmed sends.""" + + return self.actions_confirmed + + @property + def fault_reason(self) -> str | None: + """Return the primary fault without hiding a cleanup-only fault.""" + + return ( + self.primary_fault_reason + or self.audit_fault_reason + or self.cleanup_fault_reason + ) + + +@dataclass +class _CycleContext: + cycle: int = 0 + observation: RolloutObservation | None = None + current_state_deg: np.ndarray | None = None + predicted_first_action: np.ndarray | None = None + safety_result: SafetyDecision | None = None + send_result: np.ndarray | None = None + inference_latency_s: float | None = None + + +class RolloutRunner: + """Observe, infer a chunk, validate its first action, and optionally send it.""" + + def __init__( + self, + *, + policy: PolicyAdapter, + robot: RobotAdapter, + safety: SafetyGovernor, + mode: str, + log_path: Path | str, + monotonic_clock: Callable[[], float], + stop_requested: Callable[[], bool] | _StopEvent, + action_guard: ActionGuard | None = None, + operator_verdict: Callable[[], str | None] | None = None, + ) -> None: + if mode not in ("shadow", "live"): + raise ValueError("Rollout mode must be exactly 'shadow' or 'live'") + if safety.mode != mode: + raise ValueError("Rollout mode must match the safety governor mode") + + self.policy = policy + self.robot = robot + self.safety = safety + self.mode = mode + self.log_path = Path(log_path) + self.monotonic_clock = monotonic_clock + self.action_guard = action_guard + self.operator_verdict = operator_verdict + self.stop_requested = ( + stop_requested if callable(stop_requested) else stop_requested.is_set + ) + + def run(self, max_cycles: int) -> RunSummary: + """Run until the cycle limit, a stop request, or a fail-closed fault.""" + + if ( + isinstance(max_cycles, bool) + or not isinstance(max_cycles, int) + or max_cycles <= 0 + ): + raise ValueError("Rollout max_cycles must be a positive integer") + return self._run(max_cycles=max_cycles, episode=None) + + def run_episode( + self, + *, + attempt: int, + timeout_s: float = EPISODE_TIMEOUT_S, + max_actions: int = EPISODE_MAX_ACTIONS, + ) -> RunSummary: + """Run one human-verdict episode with fixed fail-closed bounds.""" + + if isinstance(attempt, bool) or not isinstance(attempt, int) or attempt < 1: + raise ValueError("Episode attempt must be a positive integer") + if not 0.0 < float(timeout_s) <= EPISODE_TIMEOUT_S: + raise ValueError("Episode timeout must be within (0, 30.0] seconds") + if ( + isinstance(max_actions, bool) + or not isinstance(max_actions, int) + or not 1 <= max_actions <= EPISODE_MAX_ACTIONS + ): + raise ValueError("Episode action cap must be within [1, 300]") + if self.operator_verdict is None: + raise ValueError("Episode mode requires a nonblocking operator verdict source") + episode = (attempt, float(timeout_s), max_actions) + return self._run(max_cycles=None, episode=episode) + + def _run( + self, + *, + max_cycles: int | None, + episode: tuple[int, float, int] | None, + ) -> RunSummary: + """Execute the shared one-step loop in bounded smoke or episode mode.""" + + cycles_completed = 0 + actions_attempted = 0 + actions_confirmed = 0 + clamp_count = 0 + attempt = 1 if episode is None else episode[0] + elapsed_seconds = 0.0 + episode_started_at = ( + None if episode is None else self.monotonic_clock() + ) + terminal_reason = "max_cycles" + primary_fault_reason: str | None = None + cleanup_fault_reason: str | None = None + audit_fault_reason: str | None = None + phase = "audit log open" + context = _CycleContext() + log_file: IO[str] | None = None + + def episode_boundary() -> str | None: + nonlocal elapsed_seconds + + if episode is None or episode_started_at is None: + return None + checked_at = self.monotonic_clock() + elapsed_seconds = max(0.0, checked_at - episode_started_at) + if self.stop_requested(): + return "stopped" + verdict = self.operator_verdict() + if self.stop_requested(): + return "stopped" + if verdict is not None and verdict not in _OPERATOR_VERDICTS: + raise ValueError(f"Unsupported operator verdict: {verdict!r}") + if verdict == "success": + return "operator_success" + if verdict == "failure": + return "operator_failure" + if elapsed_seconds >= episode[1]: + return "timeout" + if self.mode == "live" and actions_confirmed >= episode[2]: + return "timeout" + return None + + def emit(event: str, *, monotonic_s: float | None = None) -> bool: + nonlocal audit_fault_reason + nonlocal primary_fault_reason + nonlocal terminal_reason + + if log_file is None: + return False + error = self._write_event_guarded( + log_file, + event=event, + context=context, + actions_attempted=actions_attempted, + actions_confirmed=actions_confirmed, + attempt=attempt, + elapsed_seconds=elapsed_seconds, + clamp_count=clamp_count, + terminal_reason=(terminal_reason if event == "terminal" else None), + primary_fault_reason=primary_fault_reason, + cleanup_fault_reason=cleanup_fault_reason, + audit_fault_reason=audit_fault_reason, + monotonic_s=monotonic_s, + ) + if error is None: + return True + + audit_fault_reason = f"audit log {event} failed: {error}" + if primary_fault_reason is None: + primary_fault_reason = audit_fault_reason + terminal_reason = "fault" + self._write_fallback_fault( + log_file, + failed_event=event, + context=context, + actions_attempted=actions_attempted, + actions_confirmed=actions_confirmed, + attempt=attempt, + elapsed_seconds=elapsed_seconds, + clamp_count=clamp_count, + fault_reason=audit_fault_reason, + episode_mode=episode is not None, + ) + return False + + try: + try: + self.log_path.parent.mkdir(parents=True, exist_ok=True) + log_file = self.log_path.open("a", encoding="utf-8") + except Exception as exc: + audit_fault_reason = ( + "audit log open failed: " + self._exception_text(exc) + ) + primary_fault_reason = audit_fault_reason + terminal_reason = "fault" + + if log_file is not None: + phase = "connect" + self.robot.connect() + + while max_cycles is None or cycles_completed < max_cycles: + context = _CycleContext(cycle=cycles_completed) + if episode is not None: + phase = "episode boundary" + boundary_reason = episode_boundary() + if boundary_reason is not None: + terminal_reason = boundary_reason + emit(boundary_reason) + break + phase = "stop check" + if self.stop_requested(): + terminal_reason = ( + "stopped" if episode is not None else "stop_requested" + ) + emit(terminal_reason) + break + + phase = "observation" + context.observation = self.robot.observe() + try: + current_state = np.asarray( + context.observation.state_deg, dtype=float + ).copy() + except (TypeError, ValueError) as exc: + primary_fault_reason = ( + "observation state must be numeric: " + + self._exception_text(exc) + ) + terminal_reason = "fault" + observation_logged_at = self.monotonic_clock() + if emit( + "observation", monotonic_s=observation_logged_at + ): + emit("fault", monotonic_s=observation_logged_at) + break + + current_state.setflags(write=False) + context.current_state_deg = current_state + observation_logged_at = self.monotonic_clock() + if not emit( + "observation", monotonic_s=observation_logged_at + ): + break + + if episode is not None: + phase = "episode boundary" + boundary_reason = episode_boundary() + if boundary_reason is not None: + terminal_reason = boundary_reason + emit(boundary_reason) + break + phase = "stop check" + if self.stop_requested(): + terminal_reason = ( + "stopped" if episode is not None else "stop_requested" + ) + emit(terminal_reason) + break + + policy_observation = RolloutObservation( + front=context.observation.front, + side=context.observation.side, + state_deg=current_state.copy(), + task=context.observation.task, + captured_monotonic_s=( + context.observation.captured_monotonic_s + ), + ) + phase = "inference" + inference_started_at = self.monotonic_clock() + try: + raw_prediction = self.policy.predict(policy_observation) + except Exception as exc: + inference_finished_at = self.monotonic_clock() + context.inference_latency_s = max( + 0.0, inference_finished_at - inference_started_at + ) + primary_fault_reason = ( + "inference failed: " + self._exception_text(exc) + ) + terminal_reason = "fault" + emit("fault", monotonic_s=inference_finished_at) + break + + inference_finished_at = self.monotonic_clock() + context.inference_latency_s = max( + 0.0, inference_finished_at - inference_started_at + ) + + try: + prediction = np.asarray(raw_prediction, dtype=float) + except (TypeError, ValueError) as exc: + primary_fault_reason = ( + "prediction must be a numeric array: " + + self._exception_text(exc) + ) + terminal_reason = "fault" + emit("fault", monotonic_s=inference_finished_at) + break + + if prediction.ndim == 2 and prediction.shape[0] >= 1: + context.predicted_first_action = prediction[0].copy() + if not emit( + "prediction", monotonic_s=inference_finished_at + ): + break + + if ( + prediction.ndim != 2 + or prediction.shape[0] < 1 + or prediction.shape[1] != EXPECTED_ACTION_DIMENSION + ): + primary_fault_reason = ( + "prediction shape must be [steps, 7] with at least " + f"one step; received {prediction.shape}" + ) + terminal_reason = "fault" + emit("fault", monotonic_s=inference_finished_at) + break + + phase = "stop check" + if self.stop_requested(): + terminal_reason = ( + "stopped" if episode is not None else "stop_requested" + ) + emit(terminal_reason) + break + + if self.action_guard is not None: + phase = "workspace safety validation" + workspace_checked_at = self.monotonic_clock() + try: + self.action_guard.validate( + context.predicted_first_action.copy() + ) + except Exception as exc: + primary_fault_reason = ( + "workspace safety rejected action: " + + self._exception_text(exc) + ) + terminal_reason = "fault" + if emit( + "workspace_safety_fault", + monotonic_s=workspace_checked_at, + ): + emit("fault", monotonic_s=workspace_checked_at) + break + + phase = "safety validation" + safety_checked_at = self.monotonic_clock() + context.safety_result = self.safety.validate( + context.current_state_deg, + context.predicted_first_action, + safety_checked_at, + context.observation.captured_monotonic_s, + ) + if context.safety_result.clamped: + clamp_count += 1 + if not emit("safety", monotonic_s=safety_checked_at): + break + + if ( + not context.safety_result.accepted + or context.safety_result.action_deg is None + ): + primary_fault_reason = ( + "safety rejected action: " + f"{context.safety_result.reason}" + ) + terminal_reason = "fault" + emit("fault") + break + + phase = "stop check" + if self.stop_requested(): + terminal_reason = ( + "stopped" if episode is not None else "stop_requested" + ) + emit(terminal_reason) + break + + if self.mode == "live": + if episode is not None: + phase = "episode send boundary" + boundary_reason = episode_boundary() + if boundary_reason is not None: + terminal_reason = boundary_reason + emit(boundary_reason) + break + if not emit("send_intent"): + break + if episode is not None: + phase = "episode send boundary" + boundary_reason = episode_boundary() + if boundary_reason is not None: + terminal_reason = boundary_reason + emit("send_cancelled") + emit(boundary_reason) + break + if self.stop_requested(): + terminal_reason = ( + "stopped" + if episode is not None + else "stop_requested" + ) + emit("send_cancelled") + break + + phase = "send boundary safety validation" + send_checked_at = self.monotonic_clock() + context.safety_result = self.safety.validate( + context.current_state_deg, + context.safety_result.action_deg, + send_checked_at, + context.observation.captured_monotonic_s, + ) + if ( + not context.safety_result.accepted + or context.safety_result.action_deg is None + ): + emit( + "send_boundary_safety", + monotonic_s=send_checked_at, + ) + primary_fault_reason = ( + "safety rejected action at live send boundary: " + f"{context.safety_result.reason}" + ) + terminal_reason = "fault" + emit("send_cancelled") + emit("fault") + break + + if episode is not None: + phase = "episode final send boundary" + boundary_reason = episode_boundary() + if boundary_reason is not None: + terminal_reason = boundary_reason + emit( + "send_boundary_safety", + monotonic_s=send_checked_at, + ) + emit("send_cancelled") + emit(boundary_reason) + break + + actions_attempted += 1 + phase = "send" + try: + raw_send_result = self.robot.send_action( + context.safety_result.action_deg.copy() + ) + except Exception as exc: + primary_fault_reason = ( + "send failed: " + self._exception_text(exc) + ) + terminal_reason = "fault" + emit( + "send_boundary_safety", + monotonic_s=send_checked_at, + ) + emit("send_failed") + emit("fault") + break + + try: + context.send_result = np.asarray( + raw_send_result, dtype=float + ).copy() + except (TypeError, ValueError) as exc: + primary_fault_reason = ( + "send result must be numeric: " + + self._exception_text(exc) + ) + terminal_reason = "fault" + emit( + "send_boundary_safety", + monotonic_s=send_checked_at, + ) + emit("send_failed") + emit("fault") + break + + actions_confirmed += 1 + if not emit( + "send_boundary_safety", + monotonic_s=send_checked_at, + ): + break + if not emit("send_confirmed"): + break + + cycles_completed += 1 + except Exception as exc: + if primary_fault_reason is None: + primary_fault_reason = ( + f"{phase} failed: " + self._exception_text(exc) + ) + terminal_reason = "fault" + emit("fault") + finally: + try: + self.robot.disconnect() + except Exception as exc: + cleanup_fault_reason = ( + "disconnect failed: " + self._exception_text(exc) + ) + terminal_reason = "fault" + + if episode is not None and episode_started_at is not None: + try: + elapsed_seconds = max( + 0.0, self.monotonic_clock() - episode_started_at + ) + except Exception as exc: + if primary_fault_reason is None: + primary_fault_reason = ( + "episode final clock failed: " + + self._exception_text(exc) + ) + terminal_reason = "fault" + if episode is not None and ( + terminal_reason == "fault" + or primary_fault_reason is not None + or cleanup_fault_reason is not None + or audit_fault_reason is not None + ): + terminal_reason = "safety_fault" + + if log_file is not None: + context.cycle = cycles_completed + emit("terminal") + if episode is not None and terminal_reason == "fault": + terminal_reason = "safety_fault" + try: + log_file.close() + except Exception: + pass + + return RunSummary( + mode=self.mode, + cycles_completed=cycles_completed, + actions_attempted=actions_attempted, + actions_confirmed=actions_confirmed, + terminal_reason=terminal_reason, + primary_fault_reason=primary_fault_reason, + cleanup_fault_reason=cleanup_fault_reason, + audit_fault_reason=audit_fault_reason, + attempt=attempt, + elapsed_seconds=elapsed_seconds, + clamp_count=clamp_count, + ) + + def _write_event_guarded( + self, + log_file: IO[str], + **event_fields: object, + ) -> str | None: + try: + self._write_event(log_file, **event_fields) + except Exception as exc: + return self._exception_text(exc) + return None + + def _write_event( + self, + log_file: IO[str], + *, + event: str, + context: _CycleContext, + actions_attempted: int, + actions_confirmed: int, + attempt: int, + elapsed_seconds: float, + clamp_count: int, + terminal_reason: str | None, + primary_fault_reason: str | None, + cleanup_fault_reason: str | None, + audit_fault_reason: str | None, + monotonic_s: float | None = None, + ) -> None: + checked_monotonic_s = ( + self.monotonic_clock() if monotonic_s is None else monotonic_s + ) + fault_reason = ( + primary_fault_reason or audit_fault_reason or cleanup_fault_reason + ) + row = { + "timestamp_utc": self._utc_timestamp(), + "monotonic_s": float(checked_monotonic_s), + "event": event, + "mode": self.mode, + "cycle": context.cycle, + "task": ( + None if context.observation is None else context.observation.task + ), + "current_state_deg": self._json_array(context.current_state_deg), + "predicted_first_action_deg": self._json_array( + context.predicted_first_action + ), + "safety_result": self._json_safety_result(context.safety_result), + "send_result_deg": self._json_array(context.send_result), + "inference_latency_s": context.inference_latency_s, + "actions_attempted": actions_attempted, + "actions_confirmed": actions_confirmed, + "attempt": attempt, + "elapsed_seconds": elapsed_seconds, + "clamp_count": clamp_count, + "terminal_reason": terminal_reason, + "fault_reason": fault_reason, + "primary_fault_reason": primary_fault_reason, + "cleanup_fault_reason": cleanup_fault_reason, + "audit_fault_reason": audit_fault_reason, + } + log_file.write(json.dumps(row, allow_nan=False, sort_keys=True) + "\n") + log_file.flush() + + def _write_fallback_fault( + self, + log_file: IO[str], + *, + failed_event: str, + context: _CycleContext, + actions_attempted: int, + actions_confirmed: int, + attempt: int, + elapsed_seconds: float, + clamp_count: int, + fault_reason: str, + episode_mode: bool, + ) -> None: + try: + row = { + "timestamp_utc": self._utc_timestamp(), + "monotonic_s": None, + "event": ( + "terminal_fallback" + if failed_event == "terminal" + else "fault_fallback" + ), + "mode": self.mode, + "cycle": context.cycle, + "task": None, + "current_state_deg": None, + "predicted_first_action_deg": None, + "safety_result": None, + "send_result_deg": None, + "inference_latency_s": None, + "actions_attempted": actions_attempted, + "actions_confirmed": actions_confirmed, + "attempt": attempt, + "elapsed_seconds": elapsed_seconds, + "clamp_count": clamp_count, + "failed_event": failed_event, + "fault_reason": fault_reason, + "primary_fault_reason": None, + "cleanup_fault_reason": None, + "audit_fault_reason": fault_reason, + "terminal_reason": ( + ( + "safety_fault" if episode_mode else "fault" + ) + if failed_event == "terminal" + else None + ), + } + log_file.write(json.dumps(row, allow_nan=False, sort_keys=True) + "\n") + log_file.flush() + except Exception: + pass + + @staticmethod + def _exception_text(exc: Exception) -> str: + try: + return str(exc) + except Exception: + return type(exc).__name__ + + @staticmethod + def _utc_timestamp() -> str: + return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") + + @staticmethod + def _json_array(value: np.ndarray | None) -> list[float | None] | None: + if value is None: + return None + values = np.asarray(value, dtype=float).reshape(-1).tolist() + return [float(item) if np.isfinite(item) else None for item in values] + + @classmethod + def _json_safety_result( + cls, decision: SafetyDecision | None + ) -> dict[str, object] | None: + if decision is None: + return None + return { + "accepted": decision.accepted, + "action_deg": cls._json_array(decision.action_deg), + "clamped": decision.clamped, + "reason": decision.reason, + } diff --git a/p3_vlm_orchestrator/policy_rollout/workspace_guard.py b/p3_vlm_orchestrator/policy_rollout/workspace_guard.py new file mode 100644 index 0000000..112bdb5 --- /dev/null +++ b/p3_vlm_orchestrator/policy_rollout/workspace_guard.py @@ -0,0 +1,336 @@ +"""Fail-closed calibrated workspace validation for learned follower actions. + +The learned seven-dimensional action is actuated by the LeRobot follower plugin. +The P1 ``ArmConfig`` and reBot SDK are used here only to forward-kinematics check +the first six physical follower joints; the Cartesian P1 client is not an +actuation path for policy actions. +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime, timezone +import json +import math +from pathlib import Path +import sys +from typing import Any + +import numpy as np +import yaml + + +DEFAULT_MAX_CALIBRATION_AGE_S = 12 * 60 * 60 +MAX_FUTURE_SKEW_S = 5 * 60 +DEFAULT_APPROACH_HEIGHT_MM = 40.0 +POLYGON_TOLERANCE = 1e-6 +DIMENSION_TOLERANCE_MM = 1e-3 + + +class WorkspaceViolation(RuntimeError): + """Raised when a proposed physical follower action is not workspace-safe.""" + + +@dataclass(frozen=True) +class CalibratedWorkspaceGuard: + """Validate predicted tip position against a calibrated arm-frame prism.""" + + polygon_xy_mm: np.ndarray + z_min_mm: float + z_max_mm: float + fk_deg_to_xyz_mm: Callable[[np.ndarray], Sequence[float] | np.ndarray] + + @classmethod + def from_files( + cls, + *, + arm_config_path: Path | str, + workspace_config_path: Path | str, + calibration_path: Path | str, + fk_deg_to_xyz_mm: ( + Callable[[np.ndarray], Sequence[float] | np.ndarray] | None + ) = None, + current_utc: datetime | Callable[[], datetime] | None = None, + max_calibration_age_s: float = DEFAULT_MAX_CALIBRATION_AGE_S, + ) -> CalibratedWorkspaceGuard: + """Load and authenticate the geometric inputs without touching hardware.""" + + arm_path = Path(arm_config_path).expanduser().resolve() + workspace_path = Path(workspace_config_path).expanduser().resolve() + calibration_file = Path(calibration_path).expanduser().resolve() + + arm_raw = _load_mapping(arm_path, "arm config") + workspace = _load_mapping(workspace_path, "workspace config") + calibration = _load_mapping(calibration_file, "workspace calibration") + + now = current_utc() if callable(current_utc) else current_utc + if now is None: + now = datetime.now(timezone.utc) + checked_now = _aware_utc(now, "Current UTC time") + _validate_freshness( + calibration.get("updated_at"), + now=checked_now, + max_age_s=max_calibration_age_s, + ) + + affine = calibration.get("plane_to_arm") + if not isinstance(affine, Mapping): + raise ValueError("Workspace calibration affine transform is missing") + A = _finite_array(affine.get("A"), (2, 2), "affine A") + b = _finite_array(affine.get("b"), (2,), "affine b") + + zone = workspace.get("zone") + if not isinstance(zone, Mapping): + raise ValueError("Workspace zone configuration is missing") + width = _positive_finite(zone.get("width_mm"), "workspace zone width") + depth = _positive_finite(zone.get("depth_mm"), "workspace zone depth") + plane_corners = _plane_corners(calibration) + _validate_zone_dimensions(plane_corners, width, depth) + + polygon = np.asarray([A @ corner + b for corner in plane_corners], dtype=float) + _validate_convex_polygon(polygon) + polygon.setflags(write=False) + + # ArmConfig is the single source of the P1 arm safety fields. The raw + # mapping is read only for approach_height_mm because the existing + # dataclass intentionally does not expose that field yet. + from p1_arm_motion.arm_client import ArmConfig + + arm_config = ArmConfig.from_yaml(arm_path) + grasp_heights = np.asarray( + [float(value) for value in arm_config.grasp_heights_mm.values()], + dtype=float, + ) + if grasp_heights.size == 0 or not np.all(np.isfinite(grasp_heights)): + raise ValueError("Arm grasp heights must contain finite values") + safety_raw = arm_raw.get("safety") + if safety_raw is None: + safety_raw = {} + if not isinstance(safety_raw, Mapping): + raise ValueError("Arm safety config must be an object") + approach = _finite_float( + safety_raw.get("approach_height_mm", DEFAULT_APPROACH_HEIGHT_MM), + "arm approach height", + ) + z_min = float(np.min(grasp_heights)) - 10.0 + z_max = float(arm_config.transit_height_mm) + approach + if not math.isfinite(z_min) or not math.isfinite(z_max) or z_min >= z_max: + raise ValueError("Configured workspace Z range is invalid") + + fk = fk_deg_to_xyz_mm or _make_sdk_fk(arm_config) + if not callable(fk): + raise ValueError("Workspace forward kinematics adapter must be callable") + return cls( + polygon_xy_mm=polygon, + z_min_mm=z_min, + z_max_mm=z_max, + fk_deg_to_xyz_mm=fk, + ) + + def validate(self, action_deg: np.ndarray) -> None: + """Return normally only when the predicted physical action is in bounds.""" + + try: + action = np.asarray(action_deg, dtype=float).copy() + except (TypeError, ValueError) as exc: + raise WorkspaceViolation( + "Workspace action must contain seven finite physical joint values" + ) from exc + if action.shape != (7,): + raise WorkspaceViolation("Workspace action must have shape (7,)") + if not np.all(np.isfinite(action)): + raise WorkspaceViolation("Workspace action joint values must be finite") + + try: + xyz = np.asarray(self.fk_deg_to_xyz_mm(action.copy()), dtype=float) + except WorkspaceViolation: + raise + except Exception as exc: + raise WorkspaceViolation(f"Workspace forward kinematics failed: {exc}") from exc + if xyz.shape != (3,) or not np.all(np.isfinite(xyz)): + raise WorkspaceViolation( + "Workspace forward kinematics must return three finite XYZ millimeters" + ) + + if not _point_in_convex_polygon(xyz[:2], self.polygon_xy_mm): + raise WorkspaceViolation( + f"Predicted tip XY is outside calibrated workspace: {xyz[:2].tolist()}" + ) + if ( + xyz[2] < self.z_min_mm - POLYGON_TOLERANCE + or xyz[2] > self.z_max_mm + POLYGON_TOLERANCE + ): + raise WorkspaceViolation( + "Predicted tip Z is outside configured workspace range: " + f"{float(xyz[2])} not in [{self.z_min_mm}, {self.z_max_mm}] mm" + ) + + +def _load_mapping(path: Path, label: str) -> dict[str, Any]: + try: + raw_text = path.read_text(encoding="utf-8") + value = ( + json.loads(raw_text) + if path.suffix.lower() == ".json" + else yaml.safe_load(raw_text) + ) + except (OSError, UnicodeError, json.JSONDecodeError, yaml.YAMLError) as exc: + raise ValueError(f"Cannot read {label} at {path}: {exc}") from exc + if not isinstance(value, dict): + raise ValueError(f"{label.capitalize()} must be an object") + return value + + +def _aware_utc(value: object, label: str) -> datetime: + if not isinstance(value, datetime) or value.tzinfo is None: + raise ValueError(f"{label} must be a timezone-aware datetime") + return value.astimezone(timezone.utc) + + +def _validate_freshness( + timestamp: object, + *, + now: datetime, + max_age_s: float, +) -> None: + max_age = _finite_float(max_age_s, "Maximum calibration age") + if max_age < 0: + raise ValueError("Maximum calibration age must be nonnegative") + if not isinstance(timestamp, str) or not timestamp.strip(): + raise ValueError("Workspace calibration timestamp is missing") + try: + normalized = timestamp.strip().replace("Z", "+00:00") + updated_at = datetime.fromisoformat(normalized) + except ValueError as exc: + raise ValueError("Workspace calibration timestamp is not parseable") from exc + if updated_at.tzinfo is None: + raise ValueError("Workspace calibration timestamp must include a timezone") + age_s = (now - updated_at.astimezone(timezone.utc)).total_seconds() + if age_s < -MAX_FUTURE_SKEW_S: + raise ValueError("Workspace calibration timestamp is in the future") + if age_s > max_age: + raise ValueError("Workspace calibration is stale") + + +def _finite_float(value: object, label: str) -> float: + try: + result = float(value) + except (TypeError, ValueError) as exc: + raise ValueError(f"{label} must be finite") from exc + if not math.isfinite(result): + raise ValueError(f"{label} must be finite") + return result + + +def _positive_finite(value: object, label: str) -> float: + result = _finite_float(value, label) + if result <= 0: + raise ValueError(f"{label} must be positive") + return result + + +def _finite_array(value: object, shape: tuple[int, ...], label: str) -> np.ndarray: + try: + result = np.asarray(value, dtype=float) + except (TypeError, ValueError) as exc: + raise ValueError(f"Workspace calibration {label} is malformed") from exc + if result.shape != shape or not np.all(np.isfinite(result)): + raise ValueError(f"Workspace calibration {label} is malformed or nonfinite") + return result + + +def _plane_corners(calibration: Mapping[str, Any]) -> np.ndarray: + aruco = calibration.get("aruco") + if not isinstance(aruco, Mapping): + raise ValueError("Workspace calibration plane corners are missing") + ids = aruco.get("ids") + plane = aruco.get("plane_mm") + if ( + not isinstance(ids, list) + or len(ids) != 4 + or len(set(str(value) for value in ids)) != 4 + or not isinstance(plane, Mapping) + ): + raise ValueError("Workspace calibration must define four plane corners") + try: + corners = np.asarray([plane[str(marker_id)] for marker_id in ids], dtype=float) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError("Workspace calibration plane corners are malformed") from exc + if corners.shape != (4, 2) or not np.all(np.isfinite(corners)): + raise ValueError("Workspace calibration plane corners are malformed or nonfinite") + return corners + + +def _validate_zone_dimensions(corners: np.ndarray, width: float, depth: float) -> None: + edges = np.linalg.norm(np.roll(corners, -1, axis=0) - corners, axis=1) + expected = np.asarray((width, depth, width, depth), dtype=float) + swapped = np.asarray((depth, width, depth, width), dtype=float) + if not ( + np.allclose(edges, expected, rtol=1e-6, atol=DIMENSION_TOLERANCE_MM) + or np.allclose(edges, swapped, rtol=1e-6, atol=DIMENSION_TOLERANCE_MM) + ): + raise ValueError( + "Configured workspace zone dimensions are inconsistent with calibration plane corners" + ) + + +def _validate_convex_polygon(polygon: np.ndarray) -> None: + if polygon.shape != (4, 2) or not np.all(np.isfinite(polygon)): + raise ValueError("Calibrated workspace polygon is malformed") + edge_a = np.roll(polygon, -1, axis=0) - polygon + edge_b = np.roll(polygon, -2, axis=0) - np.roll(polygon, -1, axis=0) + crosses = _cross_2d(edge_a, edge_b) + if not ( + np.all(crosses > POLYGON_TOLERANCE) + or np.all(crosses < -POLYGON_TOLERANCE) + ): + raise ValueError("Calibrated workspace polygon must be nondegenerate and convex") + + +def _point_in_convex_polygon(point: np.ndarray, polygon: np.ndarray) -> bool: + edges = np.roll(polygon, -1, axis=0) - polygon + offsets = point - polygon + crosses = _cross_2d(edges, offsets) + return bool( + np.all(crosses >= -POLYGON_TOLERANCE) + or np.all(crosses <= POLYGON_TOLERANCE) + ) + + +def _cross_2d(left: np.ndarray, right: np.ndarray) -> np.ndarray: + """Return row-wise 2D cross products without NumPy's deprecated 2D cross.""" + + return left[..., 0] * right[..., 1] - left[..., 1] * right[..., 0] + + +def _make_sdk_fk(arm_config: object) -> Callable[[np.ndarray], np.ndarray]: + """Create a lazy SDK FK callable; this path never connects to the arm.""" + + sdk_repo = Path(getattr(arm_config, "sdk_repo")).expanduser().resolve() + if not sdk_repo.is_dir(): + raise FileNotFoundError(f"reBot SDK not found at {sdk_repo}") + sdk_config_path = sdk_repo / "config" / "rebotarm.yaml" + if not sdk_config_path.is_file(): + raise ValueError(f"SDK hardware config is missing: {sdk_config_path}") + sdk_config = _load_mapping(sdk_config_path, "SDK hardware config") + configured_hardware = sdk_config.get("hardware_yaml") + expected_hardware = getattr(arm_config, "hardware_yaml", None) + if configured_hardware != expected_hardware: + raise ValueError( + "SDK hardware_yaml does not match ArmConfig.hardware_yaml: " + f"{configured_hardware!r} != {expected_hardware!r}" + ) + sdk_root = str(sdk_repo) + if sdk_root not in sys.path: + sys.path.insert(0, sdk_root) + + def fk_deg_to_xyz_mm(action_deg: np.ndarray) -> np.ndarray: + # Pinocchio and the SDK remain lazy until a live validation is requested. + from reBotArm_control_py.kinematics import joint_to_pose + + physical_six_rad = np.radians(np.asarray(action_deg[:6], dtype=float)) + position_m, _ = joint_to_pose(physical_six_rad) + return np.asarray(position_m, dtype=float) * 1000.0 + + return fk_deg_to_xyz_mm diff --git a/p3_vlm_orchestrator/tests/__init__.py b/p3_vlm_orchestrator/tests/__init__.py new file mode 100644 index 0000000..11aaa62 --- /dev/null +++ b/p3_vlm_orchestrator/tests/__init__.py @@ -0,0 +1 @@ +"""Tests for the Person 4 policy rollout harness.""" diff --git a/p3_vlm_orchestrator/tests/test_episode_control.py b/p3_vlm_orchestrator/tests/test_episode_control.py new file mode 100644 index 0000000..5b9d107 --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_episode_control.py @@ -0,0 +1,358 @@ +from __future__ import annotations + +import json +from pathlib import Path +import tempfile +import threading +import unittest +from unittest.mock import patch + +import numpy as np + +from p3_vlm_orchestrator.policy_rollout.runner import RolloutRunner +from rebot_operator_kit.rollout.contracts import RolloutObservation +from rebot_operator_kit.rollout.safety import SafetyGovernor + + +TASK = "Pick up one can and place it in the taped sorting zone" +LIMITS = np.repeat(np.array([[-200.0, 200.0]]), 7, axis=0) + + +class MutableClock: + def __init__(self, value: float = 0.0) -> None: + self.value = value + + def __call__(self) -> float: + return self.value + + +class SequenceVerdicts: + def __init__(self, *values: str | None) -> None: + self.values = iter(values) + + def __call__(self) -> str | None: + return next(self.values, None) + + +class EpisodePolicy: + def __init__(self, delta: float = 0.0) -> None: + self.delta = delta + self.calls = 0 + + def predict(self, observation: RolloutObservation) -> np.ndarray: + self.calls += 1 + action = np.asarray(observation.state_deg, dtype=float) + self.delta + return np.repeat(action[None, :], 10, axis=0) + + +class EpisodeRobot: + def __init__( + self, + clock: MutableClock, + *, + after_observe=None, + ) -> None: + self.clock = clock + self.after_observe = after_observe + self.state = np.zeros(7, dtype=float) + self.connect_count = 0 + self.disconnect_count = 0 + self.observe_count = 0 + self.sent_actions: list[np.ndarray] = [] + + def connect(self) -> None: + self.connect_count += 1 + + def disconnect(self) -> None: + self.disconnect_count += 1 + + def observe(self) -> RolloutObservation: + self.observe_count += 1 + if self.after_observe is not None: + self.after_observe() + return RolloutObservation( + front=np.zeros((2, 2, 3), dtype=np.uint8), + side=np.zeros((2, 2, 3), dtype=np.uint8), + state_deg=self.state.copy(), + task=TASK, + captured_monotonic_s=self.clock(), + ) + + def send_action(self, action_deg: np.ndarray) -> np.ndarray: + action = np.asarray(action_deg, dtype=float).copy() + self.sent_actions.append(action) + self.state = action.copy() + return action + + +class TerminalFailingAudit: + def __init__(self) -> None: + self.lines: list[str] = [] + self.failed = False + + def write(self, text: str) -> int: + row = json.loads(text) + if row.get("event") == "terminal" and not self.failed: + self.failed = True + raise OSError("terminal audit failed") + self.lines.append(text) + return len(text) + + def flush(self) -> None: + return None + + def close(self) -> None: + return None + + def rows(self) -> list[dict[str, object]]: + return [json.loads(line) for line in self.lines] + + +class EpisodeControlTest(unittest.TestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.log_path = Path(temporary.name) / "episode.jsonl" + + def make_runner( + self, + *, + mode: str, + policy: EpisodePolicy, + robot: EpisodeRobot, + clock: MutableClock, + verdicts, + stop: threading.Event | None = None, + ) -> RolloutRunner: + return RolloutRunner( + policy=policy, + robot=robot, + safety=SafetyGovernor(LIMITS, mode=mode), + mode=mode, + log_path=self.log_path, + monotonic_clock=clock, + stop_requested=stop or threading.Event(), + operator_verdict=verdicts, + ) + + def test_success_before_inference_sends_nothing_further(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts("success"), + ) + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "operator_success") + self.assertEqual(policy.calls, 0) + self.assertEqual(robot.sent_actions, []) + self.assertEqual(robot.disconnect_count, 1) + + def test_failure_consumed_immediately_before_send_prevents_motion(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts(None, None, "failure"), + ) + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "operator_failure") + self.assertEqual(policy.calls, 1) + self.assertEqual(robot.sent_actions, []) + + def test_failure_arriving_during_final_safety_check_prevents_send(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts(None, None, None, None, "failure"), + ) + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "operator_failure") + self.assertEqual(summary.actions_attempted, 0) + self.assertEqual(robot.sent_actions, []) + + def test_stop_racing_success_wins_and_sends_nothing(self) -> None: + clock = MutableClock() + stop = threading.Event() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + + def success_and_stop() -> str: + stop.set() + return "success" + + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=success_and_stop, + stop=stop, + ) + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "stopped") + self.assertEqual(policy.calls, 0) + self.assertEqual(robot.sent_actions, []) + + def test_timeout_boundary_prevents_inference_and_motion_at_30_seconds(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock, after_observe=lambda: setattr(clock, "value", 30.0)) + policy = EpisodePolicy() + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts(None, None), + ) + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "timeout") + self.assertEqual(summary.elapsed_seconds, 30.0) + self.assertEqual(policy.calls, 0) + self.assertEqual(robot.sent_actions, []) + + def test_action_cap_counts_confirmed_live_sends_and_never_sends_301(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts(), + ) + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "timeout") + self.assertEqual(summary.actions_confirmed, 300) + self.assertEqual(len(robot.sent_actions), 300) + + def test_safety_fault_is_publicly_distinguishable(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy(delta=2.0) + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts(None, None), + ) + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "safety_fault") + self.assertIn("safety rejected", summary.primary_fault_reason or "") + self.assertEqual(robot.sent_actions, []) + + def test_audit_open_fault_uses_the_episode_safety_fault_reason(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts("success"), + ) + runner.log_path = self.log_path.parent + + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "safety_fault") + self.assertIn("audit log open", summary.audit_fault_reason or "") + self.assertEqual(robot.sent_actions, []) + + def test_terminal_audit_fault_keeps_exact_episode_terminal_reason(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + audit = TerminalFailingAudit() + runner = self.make_runner( + mode="live", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts("success"), + ) + + with patch.object(Path, "open", return_value=audit): + summary = runner.run_episode(attempt=1) + + self.assertEqual(summary.terminal_reason, "safety_fault") + fallback = next( + row for row in audit.rows() if row["event"] == "terminal_fallback" + ) + self.assertEqual(fallback["terminal_reason"], "safety_fault") + + def test_shadow_episode_stays_send_free(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock) + policy = EpisodePolicy() + runner = self.make_runner( + mode="shadow", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts(None, None, "success"), + ) + + summary = runner.run_episode(attempt=2) + + self.assertEqual(summary.terminal_reason, "operator_success") + self.assertEqual(policy.calls, 1) + self.assertEqual(robot.sent_actions, []) + + def test_terminal_summary_and_audit_include_attempt_elapsed_and_clamps(self) -> None: + clock = MutableClock() + robot = EpisodeRobot(clock, after_observe=lambda: setattr(clock, "value", 1.0)) + policy = EpisodePolicy(delta=2.0) + runner = self.make_runner( + mode="shadow", + policy=policy, + robot=robot, + clock=clock, + verdicts=SequenceVerdicts(None, None, "success"), + ) + + summary = runner.run_episode(attempt=2) + + self.assertEqual(summary.attempt, 2) + self.assertEqual(summary.elapsed_seconds, 1.0) + self.assertEqual(summary.clamp_count, 1) + terminal = json.loads(self.log_path.read_text().splitlines()[-1]) + self.assertEqual(terminal["event"], "terminal") + self.assertEqual(terminal["attempt"], 2) + self.assertEqual(terminal["elapsed_seconds"], 1.0) + self.assertEqual(terminal["clamp_count"], 1) + self.assertEqual(terminal["terminal_reason"], "operator_success") + + +if __name__ == "__main__": + unittest.main() diff --git a/p3_vlm_orchestrator/tests/test_keyboard_stop.py b/p3_vlm_orchestrator/tests/test_keyboard_stop.py new file mode 100644 index 0000000..e5bf5aa --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_keyboard_stop.py @@ -0,0 +1,248 @@ +from __future__ import annotations + +from io import StringIO +import threading +import unittest + +from p3_vlm_orchestrator.policy_rollout.keyboard_stop import ( + FAILURE_KEYS, + STOP_KEYS, + SUCCESS_KEYS, + KeyboardStop, +) + + +class FakeSignalAPI: + SIGINT = 2 + SIGTERM = 15 + + def __init__(self) -> None: + self.current = {self.SIGINT: "old-int", self.SIGTERM: "old-term"} + self.installs: list[tuple[int, object]] = [] + + def getsignal(self, signum: int) -> object: + return self.current[signum] + + def signal(self, signum: int, handler: object) -> object: + previous = self.current[signum] + self.current[signum] = handler + self.installs.append((signum, handler)) + return previous + + +class FakeTermios: + TCSADRAIN = 1 + + def __init__(self) -> None: + self.saved = ["saved-terminal-state"] + self.get_calls: list[int] = [] + self.set_calls: list[tuple[int, int, object]] = [] + + def tcgetattr(self, fd: int) -> object: + self.get_calls.append(fd) + return self.saved + + def tcsetattr(self, fd: int, when: int, state: object) -> None: + self.set_calls.append((fd, when, state)) + + +class FakeTTY: + def __init__(self) -> None: + self.calls: list[int] = [] + + def setcbreak(self, fd: int) -> None: + self.calls.append(fd) + + +class FakeInput: + def __init__(self, *, tty: bool, characters: str = "") -> None: + self.tty = tty + self.characters = list(characters) + + def isatty(self) -> bool: + return self.tty + + def fileno(self) -> int: + return 42 + + def read(self, count: int) -> str: + if not self.characters: + return "" + return self.characters.pop(0) + + +class ImmediateThread: + def __init__(self, *, target, name: str, daemon: bool) -> None: + self.target = target + self.name = name + self.daemon = daemon + self.started = False + self.joined = False + + def start(self) -> None: + self.started = True + self.target() + + def join(self, timeout: float | None = None) -> None: + self.joined = True + + +class PassiveThread(ImmediateThread): + def start(self) -> None: + self.started = True + + +class RaisingThreadFactory: + def __call__(self, **kwargs): + raise RuntimeError("thread construction exploded") + + +class KeyboardStopTest(unittest.TestCase): + def make_stop( + self, + *, + stdin: FakeInput, + thread_factory=PassiveThread, + is_main_thread=lambda: True, + ) -> tuple[KeyboardStop, FakeSignalAPI, FakeTermios, FakeTTY, StringIO]: + signals = FakeSignalAPI() + termios = FakeTermios() + tty = FakeTTY() + warnings = StringIO() + stop = KeyboardStop( + stdin=stdin, + warning_stream=warnings, + signal_api=signals, + termios_api=termios, + tty_api=tty, + select_fn=lambda readers, writers, errors, timeout: (readers, [], []), + thread_factory=thread_factory, + is_main_thread=is_main_thread, + ) + return stop, signals, termios, tty, warnings + + def test_tty_key_sets_the_single_shared_event_and_restores_everything(self) -> None: + for key in ("q", "x", "\x1b"): + with self.subTest(key=repr(key)): + stop, signals, termios, tty, _warnings = self.make_stop( + stdin=FakeInput(tty=True, characters=key), + thread_factory=ImmediateThread, + ) + + with stop as entered: + self.assertIs(entered, stop) + self.assertIsInstance(stop.event, threading.Event) + self.assertTrue(stop.event.is_set()) + + self.assertEqual(termios.get_calls, [42]) + self.assertEqual(tty.calls, [42]) + self.assertEqual(termios.set_calls, [(42, termios.TCSADRAIN, termios.saved)]) + self.assertEqual(signals.current[signals.SIGINT], "old-int") + self.assertEqual(signals.current[signals.SIGTERM], "old-term") + + def test_exact_verdict_and_stop_keys_are_nonoverlapping(self) -> None: + self.assertEqual(SUCCESS_KEYS, frozenset(("s",))) + self.assertEqual(FAILURE_KEYS, frozenset(("f",))) + self.assertEqual(STOP_KEYS, frozenset(("q", "x", "\x1b"))) + self.assertFalse((SUCCESS_KEYS | FAILURE_KEYS) & STOP_KEYS) + + def test_tty_success_and_failure_keys_publish_a_nonblocking_verdict(self) -> None: + for key, expected in (("s", "success"), ("f", "failure")): + with self.subTest(key=key): + stop, _signals, _termios, _tty, _warnings = self.make_stop( + stdin=FakeInput(tty=True, characters=key), + thread_factory=ImmediateThread, + ) + + with stop: + self.assertEqual(stop.verdict(), expected) + self.assertFalse(stop.event.is_set()) + + def test_queued_stop_after_verdict_still_sets_the_shared_stop_event(self) -> None: + for characters, expected in (("sq", "success"), ("f\x1b", "failure")): + with self.subTest(characters=repr(characters)): + stop, _signals, _termios, _tty, _warnings = self.make_stop( + stdin=FakeInput(tty=True, characters=characters), + thread_factory=ImmediateThread, + ) + + with stop: + self.assertEqual(stop.verdict(), expected) + self.assertTrue(stop.event.is_set()) + + def test_sigint_and_sigterm_handlers_set_the_same_event(self) -> None: + for signum in (FakeSignalAPI.SIGINT, FakeSignalAPI.SIGTERM): + with self.subTest(signum=signum): + stop, signals, _termios, _tty, _warnings = self.make_stop( + stdin=FakeInput(tty=True) + ) + with stop: + handler = signals.current[signum] + self.assertTrue(callable(handler)) + handler(signum, None) + self.assertTrue(stop.event.is_set()) + + def test_non_tty_keeps_signal_stop_and_warns_without_terminal_mutation(self) -> None: + stop, signals, termios, tty, warnings = self.make_stop( + stdin=FakeInput(tty=False) + ) + + with stop: + signals.current[signals.SIGTERM](signals.SIGTERM, None) + + self.assertTrue(stop.event.is_set()) + self.assertIn("non-TTY", warnings.getvalue()) + self.assertEqual(termios.get_calls, []) + self.assertEqual(tty.calls, []) + + def test_non_main_thread_fails_closed_before_installing_anything(self) -> None: + stop, signals, _termios, _tty, _warnings = self.make_stop( + stdin=FakeInput(tty=False), + is_main_thread=lambda: False, + ) + + with self.assertRaisesRegex(RuntimeError, "main thread"): + with stop: + self.fail("non-main context entered") + + self.assertEqual(signals.installs, []) + self.assertEqual(_termios.get_calls, []) + self.assertEqual(_warnings.getvalue(), "") + + def test_thread_factory_construction_failure_restores_entry_side_effects(self) -> None: + stop, signals, termios, tty, _warnings = self.make_stop( + stdin=FakeInput(tty=True), + thread_factory=RaisingThreadFactory(), + ) + + with self.assertRaisesRegex(RuntimeError, "thread construction exploded"): + stop.__enter__() + + self.assertEqual(tty.calls, [42]) + self.assertEqual( + termios.set_calls, + [(42, termios.TCSADRAIN, termios.saved)], + ) + self.assertEqual(signals.current[signals.SIGINT], "old-int") + self.assertEqual(signals.current[signals.SIGTERM], "old-term") + + def test_stop_and_context_exit_are_idempotent_and_restore_on_error(self) -> None: + stop, signals, termios, _tty, _warnings = self.make_stop( + stdin=FakeInput(tty=True) + ) + + with self.assertRaisesRegex(RuntimeError, "boom"): + with stop: + stop.stop() + stop.stop() + raise RuntimeError("boom") + + stop.__exit__(None, None, None) + self.assertTrue(stop.event.is_set()) + self.assertEqual(len(termios.set_calls), 1) + self.assertEqual(signals.current[signals.SIGINT], "old-int") + self.assertEqual(signals.current[signals.SIGTERM], "old-term") + + +if __name__ == "__main__": + unittest.main() diff --git a/p3_vlm_orchestrator/tests/test_lerobot_policy.py b/p3_vlm_orchestrator/tests/test_lerobot_policy.py new file mode 100644 index 0000000..23a303b --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_lerobot_policy.py @@ -0,0 +1,712 @@ +from __future__ import annotations + +from contextlib import contextmanager +from dataclasses import replace +import importlib +import importlib.metadata +from pathlib import Path +from io import StringIO +import json +import sys +import tempfile +from types import ModuleType, SimpleNamespace +import unittest +from unittest.mock import patch + +import numpy as np + +from p3_vlm_orchestrator.policy_rollout.lerobot_policy import ( + LeRobotCompatibilityError, + LeRobotPolicyAdapter, + _LeRobotAPI, + _import_lerobot_api, +) +from p3_vlm_orchestrator.policy_rollout.offline import evaluate_checkpoint +from rebot_operator_kit.rollout.checkpoint import CheckpointBundle + + +TASK = "Pick up one can and place it in the taped sorting zone" + + +def checkpoint_bundle(path: Path) -> CheckpointBundle: + return CheckpointBundle( + path=path, + task=TASK, + action_dimension=7, + chunk_size=10, + action_steps=10, + image_order=( + "observation.images.front", + "observation.images.side", + ), + profile_digest="test-digest", + profile_snapshot={}, + ) + + +class ImportBoundaryTest(unittest.TestCase): + def test_uses_lerobot_registrar_only_for_third_party_policy_plugins(self) -> None: + imported_plugins: list[str] = [] + + def package(name: str) -> ModuleType: + module = ModuleType(name) + module.__path__ = [] # type: ignore[attr-defined] + return module + + torch = ModuleType("torch") + torch.inference_mode = lambda: None # type: ignore[attr-defined] + configs_policies = ModuleType("lerobot.configs.policies") + configs_policies.PreTrainedConfig = type( # type: ignore[attr-defined] + "PreTrainedConfig", + (), + {"from_pretrained": classmethod(lambda cls, path, **kwargs: None)}, + ) + policies_factory = ModuleType("lerobot.policies.factory") + policies_factory.get_policy_class = lambda name: object # type: ignore[attr-defined] + policies_factory.make_pre_post_processors = lambda **kwargs: ( # type: ignore[attr-defined] + object(), + object(), + ) + policies_utils = ModuleType("lerobot.policies.utils") + policies_utils.prepare_observation_for_inference = ( # type: ignore[attr-defined] + lambda observation, device, task: observation + ) + import_utils = ModuleType("lerobot.utils.import_utils") + + def register_third_party_plugins() -> None: + for distribution in importlib.metadata.distributions(): + name = distribution.metadata.get("Name") + if isinstance(name, str) and name.startswith( + ( + "lerobot_robot_", + "lerobot_camera_", + "lerobot_teleoperator_", + "lerobot_policy_", + ) + ): + importlib.import_module(name) + + import_utils.register_third_party_plugins = ( # type: ignore[attr-defined] + register_third_party_plugins + ) + modules = { + "torch": torch, + "lerobot": package("lerobot"), + "lerobot.configs": package("lerobot.configs"), + "lerobot.configs.policies": configs_policies, + "lerobot.policies": package("lerobot.policies"), + "lerobot.policies.factory": policies_factory, + "lerobot.policies.utils": policies_utils, + "lerobot.utils": package("lerobot.utils"), + "lerobot.utils.import_utils": import_utils, + } + distributions = [ + SimpleNamespace(metadata={"Name": "lerobot_robot_hardware"}), + SimpleNamespace(metadata={"Name": "lerobot_teleoperator_hardware"}), + SimpleNamespace(metadata={"Name": "lerobot_policy_custom"}), + ] + + with ( + patch.dict(sys.modules, modules), + patch.object( + importlib.metadata, + "distributions", + return_value=distributions, + ), + patch.object(importlib.metadata, "version", return_value="99.0"), + patch.object( + importlib, + "import_module", + side_effect=lambda name: imported_plugins.append(name) + or ModuleType(name), + ), + ): + _import_lerobot_api() + + self.assertEqual(imported_plugins, ["lerobot_policy_custom"]) + + +class FakeConfig: + def __init__(self) -> None: + self.type = "registered_test_policy" + self.device = "training-device" + + +class FakePolicy: + def __init__(self) -> None: + self.to_calls: list[str] = [] + self.eval_calls = 0 + self.reset_calls = 0 + self.prediction: object | None = None + self.predict_calls: list[object] = [] + + def to(self, device: str) -> FakePolicy: + self.to_calls.append(device) + return self + + def eval(self) -> FakePolicy: + self.eval_calls += 1 + return self + + def reset(self) -> None: + self.reset_calls += 1 + + def predict_action_chunk(self, observation: object) -> object: + self.predict_calls.append(observation) + return self.prediction + + +class FakePolicyClass: + calls: list[tuple[str, FakeConfig, bool]] = [] + policy = FakePolicy() + + @classmethod + def from_pretrained( + cls, + checkpoint: str, + *, + config: FakeConfig, + local_files_only: bool, + ) -> FakePolicy: + cls.calls.append((checkpoint, config, local_files_only)) + return cls.policy + + +class FakeProcessor: + def __init__(self) -> None: + self.reset_calls = 0 + self.reset_error: Exception | None = None + self.calls: list[object] = [] + self.transform = lambda value: value + + def reset(self) -> None: + self.reset_calls += 1 + if self.reset_error is not None: + raise self.reset_error + + def __call__(self, value: object) -> object: + self.calls.append(value) + return self.transform(value) + + +class FakeTensor: + def __init__(self, array: np.ndarray) -> None: + self.array = array + self.detach_calls = 0 + self.cpu_calls = 0 + + def detach(self) -> FakeTensor: + self.detach_calls += 1 + return self + + def cpu(self) -> FakeTensor: + self.cpu_calls += 1 + return self + + def numpy(self) -> np.ndarray: + return self.array + + def item(self) -> object: + return self.array.item() + + +class FakeBackendFixture: + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.checkpoint = Path(temporary.name).resolve() + self.config = FakeConfig() + self.preprocessor = FakeProcessor() + self.postprocessor = FakeProcessor() + self.config_calls: list[str] = [] + self.config_local_only_calls: list[bool] = [] + self.policy_type_calls: list[str] = [] + self.processor_calls: list[dict[str, object]] = [] + self.prepare_calls: list[tuple[dict[str, np.ndarray], str, str]] = [] + self.mutate_during_prepare = False + self.prepared_task_override: str | None = None + self.inference_entries = 0 + FakePolicyClass.calls = [] + FakePolicyClass.policy = FakePolicy() + + def config_from_pretrained( + path: str, *, local_files_only: bool + ) -> FakeConfig: + self.config_calls.append(path) + self.config_local_only_calls.append(local_files_only) + return self.config + + def get_policy_class(policy_type: str) -> type[FakePolicyClass]: + self.policy_type_calls.append(policy_type) + return FakePolicyClass + + def make_pre_post_processors(**kwargs: object) -> tuple[FakeProcessor, FakeProcessor]: + self.processor_calls.append(kwargs) + return self.preprocessor, self.postprocessor + + def prepare_observation( + observation: dict[str, np.ndarray], device: str, task: str + ) -> dict[str, object]: + self.prepare_calls.append((dict(observation), device, task)) + if self.mutate_during_prepare: + for value in observation.values(): + value[...] = 99 + observation["task"] = ( + task + if self.prepared_task_override is None + else self.prepared_task_override + ) + observation["robot_type"] = "" + return observation + + @contextmanager + def inference_mode(): + self.inference_entries += 1 + yield + + self.api = _LeRobotAPI( + config_from_pretrained=config_from_pretrained, + get_policy_class=get_policy_class, + make_pre_post_processors=make_pre_post_processors, + prepare_observation=prepare_observation, + inference_mode=inference_mode, + runtime_description="fake LeRobot 99 on Python 99", + ) + + +class CheckpointLoadingTest(FakeBackendFixture, unittest.TestCase): + def test_loads_config_weights_and_saved_processors_from_one_directory(self) -> None: + with patch( + "p3_vlm_orchestrator.policy_rollout.lerobot_policy._import_lerobot_api", + return_value=self.api, + ): + adapter = LeRobotPolicyAdapter.from_checkpoint( + checkpoint_bundle(self.checkpoint), "cpu" + ) + + checkpoint = str(self.checkpoint) + self.assertEqual(self.config_calls, [checkpoint]) + self.assertEqual(self.config_local_only_calls, [True]) + self.assertEqual(self.policy_type_calls, ["registered_test_policy"]) + self.assertEqual(FakePolicyClass.calls, [(checkpoint, self.config, True)]) + self.assertEqual(len(self.processor_calls), 1) + self.assertIs(self.processor_calls[0]["policy_cfg"], self.config) + self.assertEqual(self.processor_calls[0]["pretrained_path"], checkpoint) + self.assertEqual( + self.processor_calls[0]["preprocessor_config_filename"], + "preprocessor_config.json", + ) + self.assertEqual( + self.processor_calls[0]["postprocessor_config_filename"], + "postprocessor_config.json", + ) + self.assertIs(adapter.policy, FakePolicyClass.policy) + + def test_sets_requested_device_and_resets_loaded_inference_components(self) -> None: + with patch( + "p3_vlm_orchestrator.policy_rollout.lerobot_policy._import_lerobot_api", + return_value=self.api, + ): + LeRobotPolicyAdapter.from_checkpoint( + checkpoint_bundle(self.checkpoint), "mps" + ) + + self.assertEqual(self.config.device, "mps") + self.assertEqual(FakePolicyClass.policy.to_calls, ["mps"]) + self.assertEqual(FakePolicyClass.policy.eval_calls, 1) + self.assertEqual(FakePolicyClass.policy.reset_calls, 1) + self.assertEqual(self.preprocessor.reset_calls, 1) + self.assertEqual(self.postprocessor.reset_calls, 1) + + def test_names_the_unresolved_registered_policy_in_compatibility_error(self) -> None: + self.config.type = "molmoact2" + + def unavailable(policy_type: str) -> type: + raise ValueError(f"Policy type {policy_type!r} is not installed") + + api = replace(self.api, get_policy_class=unavailable) + with patch( + "p3_vlm_orchestrator.policy_rollout.lerobot_policy._import_lerobot_api", + return_value=api, + ): + with self.assertRaisesRegex( + LeRobotCompatibilityError, + r"registered policy 'molmoact2' resolution.*fake LeRobot 99", + ): + LeRobotPolicyAdapter.from_checkpoint( + checkpoint_bundle(self.checkpoint), "cpu" + ) + + def test_names_both_saved_processor_files_in_schema_error(self) -> None: + def incompatible_processors(**kwargs: object) -> tuple[object, object]: + raise KeyError("unknown processor step from newer checkpoint") + + api = replace( + self.api, + make_pre_post_processors=incompatible_processors, + ) + with patch( + "p3_vlm_orchestrator.policy_rollout.lerobot_policy._import_lerobot_api", + return_value=api, + ): + with self.assertRaisesRegex( + LeRobotCompatibilityError, + r"saved processor schema loading .*preprocessor_config.json.*postprocessor_config.json", + ): + LeRobotPolicyAdapter.from_checkpoint( + checkpoint_bundle(self.checkpoint), "cpu" + ) + + def test_missing_saved_processor_state_fails_before_factory_loading(self) -> None: + (self.checkpoint / "preprocessor_config.json").write_text( + json.dumps( + { + "name": "policy_preprocessor", + "steps": [ + { + "registry_name": "normalizer_processor", + "config": {"device": "cuda"}, + "state_file": "missing-normalizer.safetensors", + } + ], + } + ) + ) + (self.checkpoint / "postprocessor_config.json").write_text( + json.dumps({"name": "policy_postprocessor", "steps": []}) + ) + + with patch( + "p3_vlm_orchestrator.policy_rollout.lerobot_policy._import_lerobot_api", + return_value=self.api, + ): + with self.assertRaisesRegex( + LeRobotCompatibilityError, + r"saved processor state file is missing.*missing-normalizer.safetensors", + ): + LeRobotPolicyAdapter.from_checkpoint( + checkpoint_bundle(self.checkpoint), "cpu" + ) + + self.assertEqual(self.processor_calls, []) + + +class AdapterInferenceTest(FakeBackendFixture, unittest.TestCase): + def make_adapter(self) -> LeRobotPolicyAdapter: + with patch( + "p3_vlm_orchestrator.policy_rollout.lerobot_policy._import_lerobot_api", + return_value=self.api, + ): + return LeRobotPolicyAdapter.from_checkpoint( + checkpoint_bundle(self.checkpoint), "cpu" + ) + + def observation(self) -> object: + from rebot_operator_kit.rollout.contracts import RolloutObservation + + return RolloutObservation( + front=np.full((3, 4, 3), 11, dtype=np.uint8), + side=np.full((2, 5, 3), 22, dtype=np.uint8), + state_deg=np.arange(7, dtype=np.float32), + task="caller task must not replace locked task", + captured_monotonic_s=123.0, + ) + + def test_reset_clears_policy_and_saved_processor_state_and_fails_closed(self) -> None: + adapter = self.make_adapter() + + adapter.reset() + + self.assertEqual(FakePolicyClass.policy.reset_calls, 2) + self.assertEqual(self.preprocessor.reset_calls, 2) + self.assertEqual(self.postprocessor.reset_calls, 2) + + self.preprocessor.reset_error = RuntimeError("state reset exploded") + with self.assertRaisesRegex(RuntimeError, "preprocessor reset failed"): + adapter.reset() + + def test_maps_raw_observation_predicts_and_postprocesses_a_chunk(self) -> None: + adapter = self.make_adapter() + observation = self.observation() + preprocessed = {"complete": "preprocessed batch"} + raw_prediction = FakeTensor(np.full((1, 2, 7), -0.25, dtype=np.float32)) + postprocessed_array = np.arange(14, dtype=np.float32).reshape(1, 2, 7) + postprocessed = FakeTensor(postprocessed_array) + self.preprocessor.transform = lambda value: preprocessed + FakePolicyClass.policy.prediction = raw_prediction + self.postprocessor.transform = lambda value: postprocessed + + result = adapter.predict(observation) + + self.assertEqual(len(self.prepare_calls), 1) + raw, device, task = self.prepare_calls[0] + self.assertEqual( + list(raw), + [ + "observation.images.front", + "observation.images.side", + "observation.state", + ], + ) + self.assertIsNot(raw["observation.images.front"], observation.front) + self.assertIsNot(raw["observation.images.side"], observation.side) + self.assertIsNot(raw["observation.state"], observation.state_deg) + self.assertEqual(device, "cpu") + self.assertEqual(task, TASK) + prepared = self.preprocessor.calls[0] + self.assertEqual( + list(prepared), + [ + "observation.images.front", + "observation.images.side", + "observation.state", + "task", + ], + ) + self.assertEqual(prepared["task"], TASK) + self.assertEqual(FakePolicyClass.policy.predict_calls, [preprocessed]) + self.assertEqual(self.postprocessor.calls, [raw_prediction]) + self.assertEqual(self.inference_entries, 1) + self.assertEqual(result.shape, (2, 7)) + self.assertEqual(result.dtype, np.float64) + np.testing.assert_array_equal(result, np.arange(14).reshape(2, 7)) + postprocessed_array[:] = 999 + self.assertFalse(np.any(result == 999)) + self.assertEqual(postprocessed.detach_calls, 1) + self.assertEqual(postprocessed.cpu_calls, 1) + + def test_copies_images_and_coerces_copied_state_before_mutating_prepare(self) -> None: + adapter = self.make_adapter() + observation = self.observation() + object.__setattr__( + observation, + "state_deg", + np.arange(7, dtype=np.float64), + ) + original_front = observation.front.copy() + original_side = observation.side.copy() + original_state = observation.state_deg.copy() + self.mutate_during_prepare = True + FakePolicyClass.policy.prediction = FakeTensor( + np.zeros((1, 2, 7), dtype=np.float32) + ) + + result = adapter.predict(observation) + + raw, _, _ = self.prepare_calls[0] + self.assertIsNot(raw["observation.images.front"], observation.front) + self.assertIsNot(raw["observation.images.side"], observation.side) + self.assertIsNot(raw["observation.state"], observation.state_deg) + self.assertEqual(raw["observation.state"].dtype, np.float32) + np.testing.assert_array_equal(observation.front, original_front) + np.testing.assert_array_equal(observation.side, original_side) + np.testing.assert_array_equal(observation.state_deg, original_state) + self.assertEqual(observation.state_deg.dtype, np.float64) + self.assertEqual(result.dtype, np.float64) + self.assertEqual(result.shape, (2, 7)) + + def test_rejects_when_preparation_does_not_preserve_locked_task(self) -> None: + adapter = self.make_adapter() + self.prepared_task_override = "a different prepared task" + FakePolicyClass.policy.prediction = FakeTensor( + np.zeros((1, 2, 7), dtype=np.float32) + ) + + with self.assertRaisesRegex(ValueError, "prepared task.*checkpoint task"): + adapter.predict(self.observation()) + + self.assertEqual(self.preprocessor.calls, []) + + def test_rejects_a_non_seven_element_state_before_preprocessing(self) -> None: + adapter = self.make_adapter() + observation = self.observation() + object.__setattr__(observation, "state_deg", np.zeros(6)) + + with self.assertRaisesRegex(ValueError, r"observation state.*\(7,\)"): + adapter.predict(observation) + + self.assertEqual(self.prepare_calls, []) + + def test_rejects_malformed_or_nonfinite_postprocessed_chunks(self) -> None: + adapter = self.make_adapter() + FakePolicyClass.policy.prediction = FakeTensor(np.zeros((1, 2, 7))) + cases = ( + (np.zeros(7), "shape"), + (np.zeros((1, 0, 7)), "at least one"), + (np.zeros((1, 2, 6)), "shape"), + (np.zeros((2, 2, 7)), "batch"), + (np.full((1, 2, 7), np.nan), "nonfinite"), + ) + + for output, message in cases: + with self.subTest(shape=output.shape, message=message): + self.postprocessor.transform = lambda value, output=output: FakeTensor( + output + ) + with self.assertRaisesRegex(ValueError, message): + adapter.predict(self.observation()) + + +class FakeDataset: + def __init__(self, samples: list[dict[str, object]]) -> None: + self.samples = samples + + def __len__(self) -> int: + return len(self.samples) + + def __getitem__(self, index: int) -> dict[str, object]: + return self.samples[index] + + +class FakeOfflineAdapter: + def __init__(self) -> None: + self.observations: list[object] = [] + self.reset_calls = 0 + self.state = 0 + + def reset(self) -> None: + self.reset_calls += 1 + self.state = 0 + + def predict(self, observation: object) -> np.ndarray: + self.observations.append(observation) + self.state += 1 + return np.full((2, 7), self.state, dtype=np.float64) + + +def dataset_sample( + episode_index: int, + *, + value: float, + task: str = TASK, +) -> dict[str, object]: + return { + "episode_index": FakeTensor(np.array(episode_index)), + "timestamp": FakeTensor(np.array(value / 10.0)), + "observation.images.front": FakeTensor( + np.full((3, 2, 5), value / 255.0, dtype=np.float32) + ), + "observation.images.side": FakeTensor( + np.full((3, 4, 6), (value + 1) / 255.0, dtype=np.float32) + ), + "observation.state": FakeTensor( + np.arange(7, dtype=np.float32) + value + ), + "task": task, + } + + +class OfflineEvaluationTest(unittest.TestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.dataset_root = Path(temporary.name).resolve() + self.bundle = checkpoint_bundle(self.dataset_root / "checkpoint") + self.adapter = FakeOfflineAdapter() + + def test_evaluates_every_sample_from_exactly_requested_distinct_episodes(self) -> None: + dataset = FakeDataset( + [ + dataset_sample(4, value=10), + dataset_sample(4, value=20), + dataset_sample(9, value=30), + dataset_sample(12, value=40), + ] + ) + loader_calls: list[Path] = [] + factory_calls: list[tuple[CheckpointBundle, str]] = [] + clock_values = iter([1.0, 1.1, 2.0, 2.2, 3.0, 3.3]) + output = StringIO() + + results = evaluate_checkpoint( + self.bundle, + self.dataset_root, + episodes=2, + device="cpu", + dataset_loader=lambda root: loader_calls.append(root) or dataset, + adapter_factory=lambda bundle, device: factory_calls.append( + (bundle, device) + ) + or self.adapter, + clock=lambda: next(clock_values), + output=output, + ) + + self.assertEqual(loader_calls, [self.dataset_root]) + self.assertEqual(factory_calls, [(self.bundle, "cpu")]) + self.assertEqual([result.episode_index for result in results], [4, 4, 9]) + self.assertEqual([result.sample_index for result in results], [0, 1, 2]) + self.assertEqual([result.shape for result in results], [(2, 7)] * 3) + self.assertEqual([result.minimum for result in results], [1.0, 2.0, 1.0]) + self.assertEqual([result.maximum for result in results], [1.0, 2.0, 1.0]) + self.assertEqual(self.adapter.reset_calls, 2) + np.testing.assert_allclose( + [result.latency_s for result in results], [0.1, 0.2, 0.3] + ) + self.assertEqual(len(self.adapter.observations), 3) + first = self.adapter.observations[0] + self.assertEqual(first.front.shape, (2, 5, 3)) + self.assertEqual(first.front.dtype, np.uint8) + self.assertTrue(np.all(first.front == 10)) + self.assertEqual(first.side.shape, (4, 6, 3)) + self.assertTrue(np.all(first.side == 11)) + np.testing.assert_array_equal(first.state_deg, np.arange(7) + 10) + self.assertEqual(first.state_deg.dtype, np.float32) + self.assertEqual(first.task, TASK) + self.assertEqual(first.captured_monotonic_s, 1.0) + printed = output.getvalue() + self.assertIn("episode=4 sample=0 shape=(2, 7)", printed) + self.assertIn("min=1.000000 max=1.000000 latency_ms=100.000", printed) + + def test_rejects_a_recorded_task_mismatch_before_prediction(self) -> None: + dataset = FakeDataset( + [dataset_sample(0, value=10, task="a different manipulation task")] + ) + + with self.assertRaisesRegex(ValueError, "dataset task mismatch"): + evaluate_checkpoint( + self.bundle, + self.dataset_root, + episodes=1, + dataset_loader=lambda root: dataset, + adapter_factory=lambda bundle, device: self.adapter, + output=StringIO(), + ) + + self.assertEqual(self.adapter.observations, []) + + def test_rejects_when_dataset_has_fewer_distinct_episodes_than_requested(self) -> None: + dataset = FakeDataset( + [ + dataset_sample(4, value=10), + dataset_sample(4, value=20), + ] + ) + + with self.assertRaisesRegex( + ValueError, + r"requested 2 distinct episodes.*found 1", + ): + evaluate_checkpoint( + self.bundle, + self.dataset_root, + episodes=2, + dataset_loader=lambda root: dataset, + adapter_factory=lambda bundle, device: self.adapter, + output=StringIO(), + ) + + def test_requires_a_positive_episode_count(self) -> None: + with self.assertRaisesRegex(ValueError, "episodes must be a positive integer"): + evaluate_checkpoint( + self.bundle, + self.dataset_root, + episodes=0, + dataset_loader=lambda root: FakeDataset([]), + adapter_factory=lambda bundle, device: self.adapter, + output=StringIO(), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/p3_vlm_orchestrator/tests/test_policy_evaluation.py b/p3_vlm_orchestrator/tests/test_policy_evaluation.py new file mode 100644 index 0000000..fde14ed --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_policy_evaluation.py @@ -0,0 +1,728 @@ +from __future__ import annotations + +from io import StringIO +import hashlib +import json +from pathlib import Path +import subprocess +import sys +import tempfile +import unittest + +try: + from p3_vlm_orchestrator.policy_rollout import evaluation +except ImportError: + evaluation = None # type: ignore[assignment] + +from p3_vlm_orchestrator.policy_rollout.cli import CliDependencies, build_parser, main + + +DIGEST_A = "a" * 64 +DIGEST_B = "b" * 64 +SUMMARY_FIELDS = ( + "checkpoint", + "trials", + "grasp_successes", + "placement_successes", + "safety_faults", + "clamps", + "mean_completion_s", + "overall_success_rate", +) + + +def make_trial( + index: int, + *, + checkpoint: str = "/models/checkpoint-a", + digest: str = DIGEST_A, + grasp_success: bool = True, + placement_success: bool = True, + terminal_reason: str = "operator_success", + safety_faults: int = 0, + clamps: int = 0, + completion_s: float | None = None, +) -> dict[str, object]: + return { + "checkpoint": checkpoint, + "checkpoint_digest": digest, + "placement_id": f"held-out-{index:02d}", + "attempts_used": 1, + "grasp_success": grasp_success, + "placement_success": placement_success, + "terminal_reason": terminal_reason, + "safety_faults": safety_faults, + "clamps": clamps, + "completion_s": float(index if completion_s is None else completion_s), + "source_jsonl_paths": [f"/audit/held-out-{index:02d}.jsonl"], + } + + +def make_manifest( + count: int = 10, + *, + checkpoint: str = "/models/checkpoint-a", + digest: str = DIGEST_A, +) -> dict[str, object]: + return { + "schema_version": 1, + "checkpoint": checkpoint, + "checkpoint_digest": digest, + "trials": [ + make_trial(index, checkpoint=checkpoint, digest=digest) + for index in range(1, count + 1) + ], + } + + +class PolicyEvaluationTest(unittest.TestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.root = Path(temporary.name) + + def require_module(self): + self.assertIsNotNone(evaluation, "evaluation module is missing") + return evaluation + + def write_manifest( + self, + manifest: dict[str, object], + name: str = "trials.json", + *, + materialize_audits: bool = True, + ) -> Path: + if materialize_audits: + audit_root = self.root / "audit" / Path(name).stem + audit_root.mkdir(parents=True, exist_ok=True) + for trial_index, trial in enumerate( + manifest.get("trials", []), # type: ignore[union-attr] + start=1, + ): + if not isinstance(trial, dict): + continue + sources = trial.get("source_jsonl_paths") + if ( + not isinstance(sources, list) + or not sources + or not all( + isinstance(source, str) and source.startswith("/audit/") + for source in sources + ) + ): + continue + audit_path = audit_root / ( + f"{trial_index:02d}-{trial.get('placement_id', 'invalid')}.jsonl" + ) + attempts = trial.get("attempts_used", 1) + attempts = attempts if isinstance(attempts, int) and not isinstance(attempts, bool) else 1 + completion = trial.get("completion_s", 0.0) + completion = ( + float(completion) + if isinstance(completion, (int, float)) + and not isinstance(completion, bool) + else 0.0 + ) + clamps = trial.get("clamps", 0) + clamps = clamps if isinstance(clamps, int) and not isinstance(clamps, bool) else 0 + rows: list[dict[str, object]] = [{ + "event": "rollout_metadata", + "checkpoint": trial.get("checkpoint"), + "checkpoint_digest": trial.get("checkpoint_digest"), + "profile_authentication": "checkpoint-sidecar-verified", + }] + for attempt in range(1, max(1, attempts) + 1): + final = attempt == attempts + rows.append( + { + "event": "terminal", + "attempt": attempt, + "terminal_reason": ( + trial.get("terminal_reason") + if final + else "operator_failure" + ), + "clamp_count": clamps if final else 0, + "elapsed_seconds": completion if final else 0.0, + } + ) + audit_path.write_text( + "".join(json.dumps(row) + "\n" for row in rows), + encoding="utf-8", + ) + trial["source_jsonl_paths"] = [str(audit_path.resolve())] + path = self.root / name + path.write_text(json.dumps(manifest), encoding="utf-8") + return path + + def test_accepts_exactly_ten_or_fifteen_distinct_final_trials(self) -> None: + module = self.require_module() + for count in (10, 15): + with self.subTest(count=count): + manifest = module.load_trial_manifest( + self.write_manifest(make_manifest(count), f"trials-{count}.json") + ) + self.assertEqual(len(manifest.trials), count) + + def test_rejects_nine_or_sixteen_trials_and_duplicate_placements(self) -> None: + module = self.require_module() + for count in (9, 16): + with self.subTest(count=count): + path = self.write_manifest(make_manifest(count), f"bad-{count}.json") + with self.assertRaisesRegex(ValueError, "10 to 15"): + module.load_trial_manifest(path) + duplicate = make_manifest() + duplicate["trials"][1]["placement_id"] = "held-out-01" # type: ignore[index] + with self.assertRaisesRegex(ValueError, "distinct placement"): + module.load_trial_manifest(self.write_manifest(duplicate, "duplicate.json")) + + def test_rejects_mixed_checkpoint_identity_or_digest(self) -> None: + module = self.require_module() + for field, value in (("checkpoint", "/models/other"), ("checkpoint_digest", DIGEST_B)): + with self.subTest(field=field): + manifest = make_manifest() + manifest["trials"][3][field] = value # type: ignore[index] + with self.assertRaisesRegex(ValueError, "mixed checkpoint"): + module.load_trial_manifest( + self.write_manifest(manifest, f"mixed-{field}.json") + ) + + def test_rejects_nonfinite_negative_or_boolean_counts_and_bad_attempts(self) -> None: + module = self.require_module() + cases = ( + ("completion_s", -0.1), + ("completion_s", float("nan")), + ("completion_s", "1.0"), + ("safety_faults", -1), + ("clamps", True), + ("attempts_used", 0), + ("attempts_used", 3), + ) + for field, value in cases: + with self.subTest(field=field, value=value): + manifest = make_manifest() + manifest["trials"][0][field] = value # type: ignore[index] + path = self.write_manifest(manifest, f"bad-{field}-{value!s}.json") + with self.assertRaises(ValueError): + module.load_trial_manifest(path) + + def test_rejects_unknown_terminal_reason_and_inconsistent_success_labels(self) -> None: + module = self.require_module() + cases = ( + {"terminal_reason": "unknown"}, + {"grasp_success": False, "placement_success": True}, + { + "terminal_reason": "safety_fault", + "safety_faults": 1, + "grasp_success": True, + "placement_success": True, + }, + {"terminal_reason": "safety_fault", "safety_faults": 0, + "grasp_success": False, "placement_success": False}, + {"terminal_reason": "operator_success", "placement_success": False}, + ) + for index, changes in enumerate(cases): + with self.subTest(changes=changes): + manifest = make_manifest() + manifest["trials"][0].update(changes) # type: ignore[index] + with self.assertRaises(ValueError): + module.load_trial_manifest( + self.write_manifest(manifest, f"inconsistent-{index}.json") + ) + + def test_rejects_missing_extra_or_invalid_audit_fields(self) -> None: + module = self.require_module() + manifests: list[dict[str, object]] = [] + missing = make_manifest() + del missing["trials"][0]["source_jsonl_paths"] # type: ignore[index] + manifests.append(missing) + extra = make_manifest() + extra["trials"][0]["training_loss"] = 0.1 # type: ignore[index] + manifests.append(extra) + invalid_source = make_manifest() + invalid_source["trials"][0]["source_jsonl_paths"] = [] # type: ignore[index] + manifests.append(invalid_source) + for index, manifest in enumerate(manifests): + with self.subTest(index=index): + with self.assertRaisesRegex(ValueError, "schema|source JSONL"): + module.load_trial_manifest( + self.write_manifest(manifest, f"schema-{index}.json") + ) + + def test_shared_retry_jsonl_is_complete_and_input_files_stay_unchanged(self) -> None: + module = self.require_module() + manifest = make_manifest() + manifest["trials"][0].update( # type: ignore[index] + attempts_used=2, + completion_s=3.25, + clamps=2, + ) + manifest_path = self.write_manifest(manifest, "shared-retry.json") + audit_path = Path(manifest["trials"][0]["source_jsonl_paths"][0]) # type: ignore[index] + before = (manifest_path.read_bytes(), audit_path.read_bytes()) + + loaded = module.load_trial_manifest(manifest_path) + + self.assertEqual(loaded.source_path, manifest_path) + self.assertEqual(len(loaded.trials[0].source_jsonl_paths), 1) + self.assertEqual( + (manifest_path.read_bytes(), audit_path.read_bytes()), + before, + ) + + def test_rejects_source_jsonl_reused_across_placements(self) -> None: + module = self.require_module() + manifest = make_manifest() + manifest["trials"][1]["completion_s"] = 1.0 # type: ignore[index] + path = self.write_manifest(manifest, "reused-source.json") + shared = manifest["trials"][0]["source_jsonl_paths"] # type: ignore[index] + manifest["trials"][1]["source_jsonl_paths"] = shared # type: ignore[index] + path.write_text(json.dumps(manifest), encoding="utf-8") + + with self.assertRaisesRegex(ValueError, "reused across placements"): + module.load_trial_manifest(path) + + def test_checkpoint_identity_binds_resolved_path_and_model_weights(self) -> None: + module = self.require_module() + checkpoint = self.root / "checkpoint" + checkpoint.mkdir() + weights = checkpoint / "model.safetensors" + weights.write_bytes(b"test weights") + + identity = module.checkpoint_identity(checkpoint) + + self.assertEqual(identity.path, checkpoint.resolve()) + self.assertEqual(identity.digest, hashlib.sha256(b"test weights").hexdigest()) + weights.unlink() + with self.assertRaisesRegex(ValueError, "model.safetensors"): + module.checkpoint_identity(checkpoint) + + def test_rejects_audit_checkpoint_digest_or_authentication_mismatch(self) -> None: + module = self.require_module() + for field, value in ( + ("checkpoint", "/models/different"), + ("checkpoint_digest", DIGEST_B), + ("profile_authentication", "standalone-untrusted"), + ): + with self.subTest(field=field): + manifest = make_manifest() + path = self.write_manifest(manifest, f"binding-{field}.json") + audit = Path(manifest["trials"][0]["source_jsonl_paths"][0]) # type: ignore[index] + rows = [json.loads(line) for line in audit.read_text().splitlines()] + metadata = next(row for row in rows if row.get("event") == "rollout_metadata") + metadata[field] = value + audit.write_text( + "".join(json.dumps(row) + "\n" for row in rows), + encoding="utf-8", + ) + with self.assertRaisesRegex(ValueError, "checkpoint|authenticated"): + module.load_trial_manifest(path) + + def test_rejects_nonabsolute_missing_symlink_or_malformed_audit_path(self) -> None: + module = self.require_module() + cases: list[tuple[str, str]] = [] + cases.append(("relative", "relative.jsonl")) + cases.append(("missing", str((self.root / "missing.jsonl").resolve()))) + target = self.root / "target.jsonl" + target.write_text("{}\n", encoding="utf-8") + symlink = self.root / "linked.jsonl" + symlink.symlink_to(target) + cases.append(("symlink", str(symlink))) + malformed = self.root / "malformed.jsonl" + malformed.write_text("not-json\n", encoding="utf-8") + cases.append(("malformed", str(malformed))) + for name, source in cases: + with self.subTest(name=name): + manifest = make_manifest() + path = self.write_manifest(manifest, f"audit-{name}.json") + manifest["trials"][0]["source_jsonl_paths"] = [source] # type: ignore[index] + path.write_text(json.dumps(manifest), encoding="utf-8") + with self.assertRaisesRegex( + ValueError, + "absolute|regular non-symlink|JSONL", + ): + module.load_trial_manifest(path) + + def test_rejects_incomplete_duplicate_or_mismatched_terminal_attempts(self) -> None: + module = self.require_module() + cases = ( + ( + "missing", + [{"event": "terminal", "attempt": 1, "terminal_reason": "operator_failure", "clamp_count": 0, "elapsed_seconds": 0.0}], + ), + ( + "duplicate", + [ + {"event": "terminal", "attempt": 1, "terminal_reason": "operator_failure", "clamp_count": 0, "elapsed_seconds": 0.0}, + {"event": "terminal", "attempt": 1, "terminal_reason": "operator_success", "clamp_count": 0, "elapsed_seconds": 3.0}, + ], + ), + ( + "final-reason", + [ + {"event": "terminal", "attempt": 1, "terminal_reason": "operator_failure", "clamp_count": 0, "elapsed_seconds": 0.0}, + {"event": "terminal", "attempt": 2, "terminal_reason": "operator_failure", "clamp_count": 0, "elapsed_seconds": 3.0}, + ], + ), + ) + for name, rows in cases: + with self.subTest(name=name): + manifest = make_manifest() + manifest["trials"][0].update(attempts_used=2, completion_s=3.0) # type: ignore[index] + path = self.write_manifest(manifest, f"terminal-{name}.json") + audit = Path(manifest["trials"][0]["source_jsonl_paths"][0]) # type: ignore[index] + metadata = json.loads(audit.read_text().splitlines()[0]) + audit.write_text( + "".join(json.dumps(row) + "\n" for row in [metadata, *rows]), + encoding="utf-8", + ) + with self.assertRaisesRegex(ValueError, "attempt|terminal reason"): + module.load_trial_manifest(path) + + def test_rejects_terminal_clamp_safety_or_completion_aggregate_mismatch(self) -> None: + module = self.require_module() + cases = ( + ("clamps", {"clamps": 2}, {"clamp_count": 1}), + ("completion", {"completion_s": 3.0}, {"elapsed_seconds": 2.0}), + ( + "safety", + { + "terminal_reason": "safety_fault", + "grasp_success": False, + "placement_success": False, + "safety_faults": 2, + }, + {}, + ), + ) + for name, trial_changes, row_changes in cases: + with self.subTest(name=name): + manifest = make_manifest() + manifest["trials"][0].update(trial_changes) # type: ignore[index] + path = self.write_manifest(manifest, f"aggregate-{name}.json") + audit = Path(manifest["trials"][0]["source_jsonl_paths"][0]) # type: ignore[index] + rows = [json.loads(line) for line in audit.read_text().splitlines()] + terminal = next(row for row in rows if row.get("event") == "terminal") + terminal.update(row_changes) + audit.write_text( + "".join(json.dumps(row) + "\n" for row in rows), + encoding="utf-8", + ) + with self.assertRaisesRegex( + ValueError, + "clamp|safety|completion", + ): + module.load_trial_manifest(path) + + def test_accepts_terminal_fallback_as_a_safety_fault_audit_row(self) -> None: + module = self.require_module() + manifest = make_manifest() + manifest["trials"][0].update( # type: ignore[index] + terminal_reason="safety_fault", + grasp_success=False, + placement_success=False, + safety_faults=1, + ) + path = self.write_manifest(manifest, "terminal-fallback.json") + audit = Path(manifest["trials"][0]["source_jsonl_paths"][0]) # type: ignore[index] + rows = [json.loads(line) for line in audit.read_text().splitlines()] + terminal = next(row for row in rows if row.get("event") == "terminal") + terminal["event"] = "terminal_fallback" + audit.write_text( + "".join(json.dumps(row) + "\n" for row in rows), + encoding="utf-8", + ) + + loaded = module.load_trial_manifest(path) + + self.assertEqual(loaded.trials[0].terminal_reason, "safety_fault") + + def test_exact_aggregation_uses_all_final_trial_durations(self) -> None: + module = self.require_module() + manifest = make_manifest() + trials = manifest["trials"] # type: ignore[assignment] + trials[0].update( # type: ignore[index] + grasp_success=False, + placement_success=False, + terminal_reason="operator_failure", + completion_s=1.0, + clamps=2, + ) + trials[1].update( # type: ignore[index] + grasp_success=True, + placement_success=False, + terminal_reason="timeout", + completion_s=2.0, + clamps=1, + ) + trials[2].update( # type: ignore[index] + grasp_success=False, + placement_success=False, + terminal_reason="safety_fault", + safety_faults=1, + completion_s=3.0, + ) + for index, trial in enumerate(trials[3:], start=4): # type: ignore[index] + trial["completion_s"] = float(index) + + loaded = module.load_trial_manifest(self.write_manifest(manifest)) + summary = module.summarize(loaded) + + self.assertEqual(tuple(summary), SUMMARY_FIELDS) + self.assertEqual( + summary, + { + "checkpoint": "/models/checkpoint-a", + "trials": 10, + "grasp_successes": 8, + "placement_successes": 7, + "safety_faults": 1, + "clamps": 3, + "mean_completion_s": 5.5, + "overall_success_rate": 0.7, + }, + ) + + def test_comparison_requires_same_placements_and_ranks_without_training_loss(self) -> None: + module = self.require_module() + first = make_manifest(checkpoint="/models/a", digest=DIGEST_A) + second = make_manifest(checkpoint="/models/b", digest=DIGEST_B) + third = make_manifest(checkpoint="/models/c", digest="c" * 64) + for trial in first["trials"][:2]: # type: ignore[index] + trial.update(grasp_success=False, placement_success=False, + terminal_reason="operator_failure") + for trial in second["trials"][:2]: # type: ignore[index] + trial.update(grasp_success=False, placement_success=False, + terminal_reason="operator_failure", clamps=1) + for trial in third["trials"][:3]: # type: ignore[index] + trial.update(grasp_success=False, placement_success=False, + terminal_reason="operator_failure") + ranked = module.compare( + [ + module.load_trial_manifest(self.write_manifest(first, "a.json")), + module.load_trial_manifest(self.write_manifest(second, "b.json")), + module.load_trial_manifest(self.write_manifest(third, "c.json")), + ] + ) + self.assertEqual([row["checkpoint"] for row in ranked], ["/models/a", "/models/b", "/models/c"]) + + changed = make_manifest(checkpoint="/models/d", digest="d" * 64) + changed["trials"][0]["placement_id"] = "different" # type: ignore[index] + with self.assertRaisesRegex(ValueError, "same held-out placement"): + module.compare( + [ + module.load_trial_manifest(self.write_manifest(first, "a2.json")), + module.load_trial_manifest(self.write_manifest(changed, "d.json")), + ] + ) + + def test_exact_ties_are_deterministic_by_checkpoint_identity(self) -> None: + module = self.require_module() + manifests = [] + for name, digest in (("z", DIGEST_B), ("a", DIGEST_A)): + path = self.write_manifest( + make_manifest(checkpoint=f"/models/{name}", digest=digest), + f"{name}.json", + ) + manifests.append(module.load_trial_manifest(path)) + ranked = module.compare(manifests) + self.assertEqual([row["checkpoint"] for row in ranked], ["/models/a", "/models/z"]) + + def test_comparison_uses_unrounded_mean_time_for_ranking(self) -> None: + module = self.require_module() + fast = make_manifest(checkpoint="/models/z-fast", digest=DIGEST_A) + slow = make_manifest(checkpoint="/models/a-slow", digest=DIGEST_B) + for trial in fast["trials"]: # type: ignore[index] + trial["completion_s"] = 1.0000001 + for trial in slow["trials"]: # type: ignore[index] + trial["completion_s"] = 1.0000002 + loaded_fast = module.load_trial_manifest( + self.write_manifest(fast, "raw-fast.json") + ) + loaded_slow = module.load_trial_manifest( + self.write_manifest(slow, "raw-slow.json") + ) + self.assertEqual( + module.summarize(loaded_fast)["mean_completion_s"], + module.summarize(loaded_slow)["mean_completion_s"], + ) + + ranked = module.compare([loaded_slow, loaded_fast]) + + self.assertEqual( + [row["checkpoint"] for row in ranked], + ["/models/z-fast", "/models/a-slow"], + ) + + def test_comparison_rejects_duplicate_digest_at_different_paths(self) -> None: + module = self.require_module() + first = module.load_trial_manifest( + self.write_manifest( + make_manifest(checkpoint="/models/a", digest=DIGEST_A), + "digest-a.json", + ) + ) + second = module.load_trial_manifest( + self.write_manifest( + make_manifest(checkpoint="/models/b", digest=DIGEST_A), + "digest-b.json", + ) + ) + + with self.assertRaisesRegex(ValueError, "digest"): + module.compare([first, second]) + + def test_ranking_breaks_success_ties_by_safety_then_clamps_then_time(self) -> None: + module = self.require_module() + specs = ( + ("fast", "d" * 64, "operator_failure", 0, 0, 1.0), + ("slow", "e" * 64, "operator_failure", 0, 0, 2.0), + ("clamped", "f" * 64, "operator_failure", 0, 1, 0.1), + ("faulted", "1" * 64, "safety_fault", 1, 0, 0.1), + ) + manifests = [] + for name, digest, reason, faults, clamps, duration in specs: + manifest = make_manifest(checkpoint=f"/models/{name}", digest=digest) + for trial in manifest["trials"]: # type: ignore[index] + trial["completion_s"] = duration + manifest["trials"][0].update( # type: ignore[index] + grasp_success=False, + placement_success=False, + terminal_reason=reason, + safety_faults=faults, + clamps=clamps, + ) + manifests.append( + module.load_trial_manifest( + self.write_manifest(manifest, f"ranking-{name}.json") + ) + ) + + ranked = module.compare(list(reversed(manifests))) + + self.assertEqual( + [row["checkpoint"] for row in ranked], + ["/models/fast", "/models/slow", "/models/clamped", "/models/faulted"], + ) + + def test_report_writer_is_deterministic_confined_and_read_only(self) -> None: + module = self.require_module() + manifest_path = self.write_manifest(make_manifest()) + before = manifest_path.read_bytes() + loaded = module.load_trial_manifest(manifest_path) + + first = module.write_reports([loaded], output_name="checkpoint-a", repo_root=self.root) + json_bytes = first.json_path.read_bytes() + csv_bytes = first.csv_path.read_bytes() + first.json_path.unlink() + first.csv_path.unlink() + second = module.write_reports([loaded], output_name="checkpoint-a", repo_root=self.root) + + self.assertEqual(json_bytes, second.json_path.read_bytes()) + self.assertEqual(csv_bytes, second.csv_path.read_bytes()) + self.assertEqual(manifest_path.read_bytes(), before) + self.assertEqual(second.json_path.parent, (self.root / "runs/policy/reports").resolve()) + self.assertEqual(csv_bytes.decode().splitlines()[0], ",".join(SUMMARY_FIELDS)) + + def test_report_writer_rejects_traversal_existing_outputs_and_symlink_escape(self) -> None: + module = self.require_module() + loaded = module.load_trial_manifest(self.write_manifest(make_manifest())) + for name in ("../escape", "checkpoint.json", "credentials", "calibration"): + with self.subTest(name=name): + with self.assertRaises(ValueError): + module.write_reports([loaded], output_name=name, repo_root=self.root) + + module.write_reports([loaded], output_name="immutable", repo_root=self.root) + with self.assertRaisesRegex(ValueError, "already exists"): + module.write_reports([loaded], output_name="immutable", repo_root=self.root) + + other = self.root / "other" + other.mkdir() + reports = self.root / "runs" / "policy" / "reports" + for child in reports.iterdir(): + child.unlink() + reports.rmdir() + reports.symlink_to(other, target_is_directory=True) + with self.assertRaisesRegex(ValueError, "symlink"): + module.write_reports([loaded], output_name="escape", repo_root=self.root) + + def test_cli_report_and_compare_help_and_output_are_lazy(self) -> None: + module = self.require_module() + del module + help_text = build_parser().format_help() + self.assertIn("report", help_text) + self.assertIn("compare", help_text) + first = self.write_manifest(make_manifest(checkpoint="/models/a", digest=DIGEST_A), "a.json") + second = self.write_manifest(make_manifest(checkpoint="/models/b", digest=DIGEST_B), "b.json") + stdout = StringIO() + stderr = StringIO() + status = main( + ["compare", "--manifest", str(first), str(second), "--output-name", "held-out"], + dependencies=CliDependencies(stdout=stdout, stderr=stderr, repo_root=self.root), + ) + self.assertEqual(status, 0, stderr.getvalue()) + self.assertIn("report_json=", stdout.getvalue()) + self.assertIn("report_csv=", stdout.getvalue()) + self.assertIn("winner=/models/a", stdout.getvalue()) + + def test_fresh_process_report_imports_no_hardware_or_model_stack(self) -> None: + self.require_module() + manifest = self.write_manifest(make_manifest()) + repo_root = Path(__file__).resolve().parents[2] + script = f""" +import sys +from pathlib import Path +from p3_vlm_orchestrator.policy_rollout.cli import CliDependencies, main +status = main( + ['report', '--manifest', {str(manifest)!r}, '--output-name', 'fresh'], + dependencies=CliDependencies(repo_root=Path({str(self.root)!r})), +) +assert status == 0, status +forbidden = [] +for name in sys.modules: + lower = name.lower() + if (name == 'serial' or name.startswith('serial.') or name == 'cv2' + or name.startswith('cv2.') or name == 'torch' or name.startswith('torch.') + or name == 'lerobot' or name.startswith('lerobot.') + or 'rebot_robot' in lower or lower.startswith('rebotarm_control_py')): + forbidden.append(name) +assert not forbidden, forbidden +""" + completed = subprocess.run( + [sys.executable, "-c", script], cwd=repo_root, + capture_output=True, text=True, check=False, + ) + self.assertEqual(completed.returncode, 0, completed.stdout + completed.stderr) + + def test_person4_runbook_documents_exact_safe_workflow(self) -> None: + runbook = Path(__file__).resolve().parents[1] / "PERSON4_RUNBOOK.md" + self.assertTrue(runbook.is_file(), "Person 4 runbook is missing") + text = runbook.read_text(encoding="utf-8") + required = ( + "Pick up one can and place it in the taped sorting zone", + "shoulder_pan, shoulder_lift, elbow_flex, wrist_flex, wrist_yaw, wrist_roll, gripper", + "front`, then `side", + "Gate A", + "Gate B", + "Gate C", + "Gate D", + "10-15", + "--episode --retry-on-failure", + "physical e-stop", + "10-20%", + "plane_to_arm", + "no automatic home", + "releases torque", + "runs/policy/reports/", + "absolute, existing, regular, non-symlink JSONL files", + "List the shared JSONL path once", + "MolmoAct2", + "No automated test performs physical motion", + ) + for phrase in required: + with self.subTest(phrase=phrase): + self.assertIn(phrase, text) + + +if __name__ == "__main__": + unittest.main() diff --git a/p3_vlm_orchestrator/tests/test_policy_rollout_cli.py b/p3_vlm_orchestrator/tests/test_policy_rollout_cli.py new file mode 100644 index 0000000..1e1b1ba --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_policy_rollout_cli.py @@ -0,0 +1,1270 @@ +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime, timezone +from io import StringIO +import hashlib +import json +from pathlib import Path +import subprocess +import sys +from types import SimpleNamespace +import tempfile +import threading +import unittest + +import numpy as np + +from p3_vlm_orchestrator.policy_rollout.cli import ( + CliDependencies, + PreflightRobotAdapter, + build_parser, + canonical_profile_digest, + default_serial_port_is_free, + main, +) +from rebot_operator_kit.rollout.contracts import RolloutObservation + + +TASK = "Pick up one can and place it in the taped sorting zone" +NOW = datetime(2026, 7, 18, 21, 22, 23, tzinfo=timezone.utc) +MANUAL_RESET_PHRASE = "I RESET THE CAN AND CLEARED THE WORKSPACE" + + +def make_profile() -> dict[str, object]: + names = ( + "shoulder_pan", + "shoulder_lift", + "elbow_flex", + "wrist_flex", + "wrist_yaw", + "wrist_roll", + "gripper", + ) + limits = ( + (-145.0, 145.0), + (-170.0, 0.0), + (-200.0, 0.0), + (-80.0, 90.0), + (-90.0, 90.0), + (-90.0, 90.0), + (-270.0, 0.0), + ) + return { + "schema_version": 1, + "profile_id": "rebot-test-profile", + "profile_version": 1, + "collection_defaults": { + "task": TASK, + "fps": 30, + "motor_velocity": 2000.0, + "gripper_force": 0.05, + }, + "training_defaults": { + "action_dimension": 7, + "chunk_size": 10, + "n_action_steps": 10, + "state_normalization": "quantile", + "action_normalization": "quantile", + "normalize_gripper": True, + "image_order": [ + "observation.images.front", + "observation.images.side", + ], + }, + "coordinate_contract": { + "frame": "follower_degrees_after_direction_limits_and_step_cap", + "control_mode": "absolute joint pose", + "action_dimension": 7, + "joints": [ + { + "name": name, + "feature": f"{name}.pos", + "leader_to_follower_scale": -1.0 if index < 2 else 1.0, + "soft_limit_degrees": list(limits[index]), + } + for index, name in enumerate(names) + ], + }, + "camera_defaults": { + "front": { + "recording_key": "observation.images.front", + "index": 0, + "width": 640, + "height": 480, + "fps": 30, + }, + "side": { + "recording_key": "observation.images.side", + "index": 1, + "width": 1280, + "height": 720, + "fps": 30, + }, + "excluded_screen_index": 3, + "excluded_screen_name": "MacBook camera", + "minimum_measured_fps": 27, + }, + "calibration": { + "follower": { + "type": "seeed_b601_dm_follower", + "id": "follower1", + "runtime_relative_path": "calibration/follower.json", + "sha256": "a" * 64, + }, + "leader": { + "type": "rebot_arm_102_leader", + "id": "rebot_arm_102_leader", + "runtime_relative_path": "calibration/leader.json", + "sha256": "e" * 64, + }, + "follower_driver_contract": { + "runtime_relative_path": "driver/config.py", + "sha256": "b" * 64, + }, + "follower_base_implementation": { + "runtime_relative_path": "driver/base.py", + "sha256": "c" * 64, + }, + "follower_dm_implementation": { + "runtime_relative_path": "driver/dm.py", + "sha256": "d" * 64, + }, + "leader_driver_contract": { + "runtime_relative_path": "driver/leader-config.py", + "sha256": "f" * 64, + }, + "leader_implementation": { + "runtime_relative_path": "driver/leader.py", + "sha256": "1" * 64, + }, + }, + "hardware_identity": { + "follower_usb": {"vid": 0x2E88, "pid": 0x4603}, + "leader_usb": {"vid": 0x1A86, "pid": 0x7523}, + "serial_port_policy": "Discover exactly one device; never persist a path.", + }, + } + + +def make_observation( + *, + task: str = TASK, + state: object | None = None, + front_shape: tuple[int, ...] = (480, 640, 3), + side_shape: tuple[int, ...] = (720, 1280, 3), +) -> RolloutObservation: + return RolloutObservation( + front=np.zeros(front_shape, dtype=np.uint8), + side=np.zeros(side_shape, dtype=np.uint8), + state_deg=np.zeros(7) if state is None else state, + task=task, + captured_monotonic_s=10.0, + ) + + +class FakeRobot: + follower_port = "/dev/fake-follower" + + def __init__(self, observations: list[RolloutObservation] | None = None) -> None: + self.observations = observations or [make_observation(), make_observation()] + self.connect_count = 0 + self.disconnect_count = 0 + self.observe_count = 0 + self.send_count = 0 + + def connect(self) -> None: + self.connect_count += 1 + + def disconnect(self) -> None: + self.disconnect_count += 1 + + def observe(self) -> RolloutObservation: + result = self.observations[min(self.observe_count, len(self.observations) - 1)] + self.observe_count += 1 + return result + + def send_action(self, action: np.ndarray) -> np.ndarray: + self.send_count += 1 + return np.asarray(action).copy() + + +class FailingConnectRobot(FakeRobot): + def __init__(self, *, cleanup_error: Exception | None = None) -> None: + super().__init__() + self.cleanup_error = cleanup_error + + def connect(self) -> None: + self.connect_count += 1 + raise RuntimeError("connect exploded") + + def disconnect(self) -> None: + self.disconnect_count += 1 + if self.cleanup_error is not None: + raise self.cleanup_error + + +class StopOnSecondObservationRobot(FakeRobot): + def __init__(self, stop_event: threading.Event) -> None: + super().__init__() + self.stop_event = stop_event + + def observe(self) -> RolloutObservation: + observation = super().observe() + if self.observe_count == 2: + self.stop_event.set() + return observation + + +class FakeKeyboardStop: + def __init__(self) -> None: + self.event = threading.Event() + self.entered = False + self.exited = False + + def __enter__(self): + self.entered = True + return self + + def __exit__(self, exc_type, exc, traceback) -> None: + self.exited = True + + def verdict(self) -> str | None: + return None + + +@dataclass +class Summary: + cycles_completed: int = 1 + actions_attempted: int = 1 + actions_confirmed: int = 1 + terminal_reason: str = "max_cycles" + primary_fault_reason: str | None = None + cleanup_fault_reason: str | None = None + audit_fault_reason: str | None = None + + +@dataclass +class EpisodeSummary: + attempt: int + terminal_reason: str + cycles_completed: int = 1 + actions_attempted: int = 0 + actions_confirmed: int = 0 + primary_fault_reason: str | None = None + cleanup_fault_reason: str | None = None + audit_fault_reason: str | None = None + elapsed_seconds: float = 1.25 + clamp_count: int = 0 + + +class ScriptedEpisodeRunner: + def __init__(self, *, outcomes: list[str], attempts: list[int], **kwargs) -> None: + self.outcomes = outcomes + self.attempts = attempts + self.robot = kwargs["robot"] + + def run_episode(self, *, attempt: int) -> EpisodeSummary: + self.attempts.append(attempt) + self.robot.connect() + self.robot.disconnect() + outcome = self.outcomes.pop(0) + return EpisodeSummary( + attempt=attempt, + terminal_reason=outcome, + primary_fault_reason=( + "simulated safety fault" if outcome == "safety_fault" else None + ), + ) + + +class ExercisingRunner: + def __init__(self, **kwargs) -> None: + self.kwargs = kwargs + + def run(self, cycles: int) -> Summary: + robot = self.kwargs["robot"] + robot.connect() + robot.observe() # Must be fresh; preflight consumed the first observation. + robot.disconnect() + return Summary( + cycles_completed=cycles, + actions_attempted=cycles if self.kwargs["mode"] == "live" else 0, + actions_confirmed=cycles if self.kwargs["mode"] == "live" else 0, + ) + + +class ConnectThenPropagateRunner: + def __init__(self, **kwargs) -> None: + self.robot = kwargs["robot"] + + def run(self, cycles: int) -> Summary: + del cycles + self.robot.connect() + raise AssertionError("connect was expected to fail") + + +class PolicyRolloutCliTest(unittest.TestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.root = Path(temporary.name) + self.stdout = StringIO() + self.stderr = StringIO() + self.profile = make_profile() + self.bundle = SimpleNamespace( + path=self.root / "checkpoint", + task=TASK, + action_dimension=7, + chunk_size=10, + action_steps=10, + image_order=("observation.images.front", "observation.images.side"), + profile_digest=canonical_profile_digest(self.profile), + profile_snapshot=self.profile, + ) + self.bundle.path.mkdir() + (self.bundle.path / "model.safetensors").write_bytes(b"test weights") + self.write_processor_configs(self.bundle.path) + + def dependencies(self, **overrides) -> CliDependencies: + values = { + "stdout": self.stdout, + "stderr": self.stderr, + "input_fn": lambda prompt: "", + "utc_now": lambda: NOW, + "monotonic_clock": lambda: 10.0, + "checkpoint_loader": lambda path: self.bundle, + "offline_evaluator": lambda *args, **kwargs: [object()], + "policy_factory": lambda bundle, device: object(), + "dummy_policy_factory": lambda: object(), + "safety_factory": lambda profile, mode: SimpleNamespace(mode=mode), + "guard_factory": lambda **kwargs: object(), + "robot_factory": lambda **kwargs: FakeRobot(), + "runner_factory": lambda **kwargs: ExercisingRunner(**kwargs), + "keyboard_stop_factory": FakeKeyboardStop, + "serial_port_is_free": lambda port: True, + "repo_root": self.root, + } + values.update(overrides) + return CliDependencies(**values) + + @staticmethod + def write_processor_configs(checkpoint: Path) -> None: + (checkpoint / "preprocessor_config.json").write_text( + json.dumps({"steps": []}), encoding="utf-8" + ) + (checkpoint / "postprocessor_config.json").write_text( + json.dumps({"steps": []}), encoding="utf-8" + ) + + def write_profile(self, profile: dict[str, object] | None = None) -> Path: + path = self.root / "training_profile.json" + path.write_text(json.dumps(profile or self.profile), encoding="utf-8") + return path + + def write_checkpoint(self, *, digest: str | None = None) -> Path: + checkpoint = self.root / "real-checkpoint" + checkpoint.mkdir() + (checkpoint / "config.json").write_text( + json.dumps({"type": "molmoact2", "device": "cuda"}), + encoding="utf-8", + ) + self.write_processor_configs(checkpoint) + (checkpoint / "model.safetensors").write_bytes(b"not-loaded-by-inspect") + collection_contract = {"task": TASK} + (checkpoint / "rebot_training_profile.json").write_text( + json.dumps( + { + "training_profile_digest": digest + or canonical_profile_digest(self.profile), + "profile_snapshot": self.profile, + "collection_contract": collection_contract, + "collection_contract_digest": canonical_profile_digest( + collection_contract + ), + } + ), + encoding="utf-8", + ) + return checkpoint + + def rollout_args(self, command: str = "live") -> list[str]: + return [ + command, + "--checkpoint", + str(self.root / "checkpoint"), + "--runtime-root", + str(self.root / "runtime"), + "--arm-config", + str(self.root / "arm.yaml"), + "--workspace-config", + str(self.root / "workspace.yaml"), + "--workspace-calibration", + str(self.root / "calibration.json"), + "--log-path", + str(self.root / "runs" / "policy" / "rollout.jsonl"), + ] + + def test_help_documents_all_staged_gates_and_physical_estop_warning(self) -> None: + help_text = build_parser().format_help() + + for gate in ("Gate A", "Gate B", "Gate C", "Gate D"): + self.assertIn(gate, help_text) + self.assertIn("physical e-stop", help_text.lower()) + self.assertIn("offline --checkpoint CHECKPOINT", help_text) + self.assertIn("live --checkpoint CHECKPOINT --cycles 5", help_text) + self.assertIn("s = success", help_text) + self.assertIn("f = failure", help_text) + self.assertIn("q/x/Esc = stop", help_text) + self.assertIn("30.0-second", help_text) + self.assertIn("300 confirmed live actions", help_text) + self.assertIn("manual reset", help_text.lower()) + self.assertIn("--retry-on-failure", help_text) + + def test_inspect_validates_and_prints_contract_without_policy_factory(self) -> None: + checkpoint = self.write_checkpoint() + deps = CliDependencies( + stdout=self.stdout, + stderr=self.stderr, + policy_factory=lambda bundle, device: self.fail("policy weights loaded"), + robot_factory=lambda **kwargs: self.fail("robot factory called"), + ) + + status = main( + ["inspect", "--checkpoint", str(checkpoint)], dependencies=deps + ) + + self.assertEqual(status, 0) + output = self.stdout.getvalue() + self.assertIn(f"locked_task={TASK}", output) + self.assertIn("policy_type=molmoact2", output) + self.assertIn("dimension=7", output) + self.assertIn( + "joint_order=shoulder_pan,shoulder_lift,elbow_flex,wrist_flex," + "wrist_yaw,wrist_roll,gripper", + output, + ) + self.assertIn("observation.images.front,observation.images.side", output) + self.assertIn("preprocessor_config.json,postprocessor_config.json", output) + self.assertIn(canonical_profile_digest(self.profile), output) + + def test_inspect_rejects_a_mismatched_checkpoint_profile_digest(self) -> None: + checkpoint = self.write_checkpoint(digest="0" * 64) + + status = main( + ["inspect", "--checkpoint", str(checkpoint)], + dependencies=CliDependencies(stdout=self.stdout, stderr=self.stderr), + ) + + self.assertEqual(status, 2) + self.assertIn("digest does not match", self.stderr.getvalue()) + + def test_offline_delegates_directly_without_touching_hardware(self) -> None: + calls: list[tuple[object, Path, int, str]] = [] + + def evaluate(bundle, dataset, *, episodes, device, output): + calls.append((bundle, dataset, episodes, device)) + self.assertIs(output, self.stdout) + return [object()] + + deps = self.dependencies( + offline_evaluator=evaluate, + robot_factory=lambda **kwargs: self.fail("hardware factory called"), + ) + status = main( + [ + "offline", + "--checkpoint", + str(self.root / "checkpoint"), + "--dataset", + str(self.root / "dataset"), + "--episodes", + "3", + "--device", + "cpu", + ], + dependencies=deps, + ) + + self.assertEqual(status, 0) + self.assertEqual(calls, [(self.bundle, self.root / "dataset", 3, "cpu")]) + + def test_live_rejects_missing_live_gate_or_out_of_range_speed_before_hardware(self) -> None: + deps = self.dependencies( + robot_factory=lambda **kwargs: self.fail("hardware factory called") + ) + + self.assertEqual(main(self.rollout_args(), dependencies=deps), 2) + self.assertIn("--live", self.stderr.getvalue()) + + self.stderr.seek(0) + self.stderr.truncate() + args = self.rollout_args() + ["--live", "--speed-scale", "0.21"] + self.assertEqual(main(args, dependencies=deps), 2) + self.assertIn("[0.10, 0.20]", self.stderr.getvalue()) + + def test_live_uses_exact_prompts_in_order_then_preflights_once(self) -> None: + prompts: list[str] = [] + replies = iter(("I HAVE AN E-STOP OPERATOR", "WORKSPACE IS EMPTY")) + robot = FakeRobot() + keyboard = FakeKeyboardStop() + runner_calls: list[dict[str, object]] = [] + + def make_runner(**kwargs): + runner_calls.append(kwargs) + return ExercisingRunner(**kwargs) + + deps = self.dependencies( + input_fn=lambda prompt: prompts.append(prompt) or next(replies), + robot_factory=lambda **kwargs: robot, + runner_factory=make_runner, + keyboard_stop_factory=lambda: keyboard, + ) + status = main( + self.rollout_args() + + ["--live", "--cycles", "2", "--speed-scale", "0.10"], + dependencies=deps, + ) + + self.assertEqual(status, 0) + self.assertEqual( + prompts, + [ + 'Type exactly "I HAVE AN E-STOP OPERATOR": ', + 'Type exactly "WORKSPACE IS EMPTY": ', + ], + ) + self.assertEqual(robot.connect_count, 1) + self.assertEqual(robot.disconnect_count, 1) + self.assertEqual(robot.observe_count, 2) + self.assertEqual(len(runner_calls), 1) + self.assertIs(runner_calls[0]["stop_requested"], keyboard.event) + self.assertIn("actions_attempted=2", self.stdout.getvalue()) + self.assertIn("primary_fault=none", self.stdout.getvalue()) + metadata = json.loads( + (self.root / "runs" / "policy" / "rollout.jsonl") + .read_text(encoding="utf-8") + .splitlines()[0] + ) + self.assertEqual(metadata["checkpoint"], str(self.bundle.path.resolve())) + self.assertEqual( + metadata["checkpoint_digest"], + hashlib.sha256(b"test weights").hexdigest(), + ) + + def test_episode_default_has_no_retry_and_prints_attempt_metrics(self) -> None: + robots: list[FakeRobot] = [] + attempts: list[int] = [] + outcomes = ["operator_failure"] + prompts: list[str] = [] + + def input_fn(prompt: str) -> str: + prompts.append(prompt) + return ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ) + + deps = self.dependencies( + input_fn=input_fn, + robot_factory=lambda **kwargs: robots.append(FakeRobot()) or robots[-1], + runner_factory=lambda **kwargs: ScriptedEpisodeRunner( + outcomes=outcomes, + attempts=attempts, + **kwargs, + ), + ) + + status = main( + self.rollout_args() + ["--live", "--episode"], + dependencies=deps, + ) + + self.assertEqual(status, 0) + self.assertEqual(attempts, [1]) + self.assertEqual(len(robots), 1) + self.assertFalse(any("manual reset" in prompt.lower() for prompt in prompts)) + summary = self.stdout.getvalue() + self.assertIn("attempt=1", summary) + self.assertIn("elapsed_seconds=1.250", summary) + self.assertIn("clamp_count=0", summary) + self.assertIn("terminal_reason=operator_failure", summary) + + def test_episode_failure_without_exact_reset_ack_is_not_retried(self) -> None: + robots: list[FakeRobot] = [] + attempts: list[int] = [] + outcomes = ["operator_failure"] + prompts: list[str] = [] + + def input_fn(prompt: str) -> str: + prompts.append(prompt) + if "manual reset" in prompt.lower(): + return "not acknowledged" + return ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ) + + deps = self.dependencies( + input_fn=input_fn, + robot_factory=lambda **kwargs: robots.append(FakeRobot()) or robots[-1], + runner_factory=lambda **kwargs: ScriptedEpisodeRunner( + outcomes=outcomes, + attempts=attempts, + **kwargs, + ), + ) + + status = main( + self.rollout_args() + + ["--live", "--episode", "--retry-on-failure"], + dependencies=deps, + ) + + self.assertEqual(status, 0) + self.assertEqual(attempts, [1]) + self.assertEqual(len(robots), 1) + self.assertEqual( + sum("manual reset" in prompt.lower() for prompt in prompts), + 1, + ) + + def test_acknowledged_failure_retries_once_with_fresh_robot_lifecycle(self) -> None: + robots: list[FakeRobot] = [] + attempts: list[int] = [] + outcomes = ["operator_failure", "operator_success"] + + def input_fn(prompt: str) -> str: + if "manual reset" in prompt.lower(): + return MANUAL_RESET_PHRASE + return ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ) + + deps = self.dependencies( + input_fn=input_fn, + robot_factory=lambda **kwargs: robots.append(FakeRobot()) or robots[-1], + runner_factory=lambda **kwargs: ScriptedEpisodeRunner( + outcomes=outcomes, + attempts=attempts, + **kwargs, + ), + ) + + status = main( + self.rollout_args() + + ["--live", "--episode", "--retry-on-failure"], + dependencies=deps, + ) + + self.assertEqual(status, 0) + self.assertEqual(attempts, [1, 2]) + self.assertEqual(len(robots), 2) + self.assertTrue(all(robot.connect_count == 1 for robot in robots)) + self.assertTrue(all(robot.disconnect_count == 1 for robot in robots)) + self.assertIn("attempt=2", self.stdout.getvalue()) + self.assertIn("terminal_reason=operator_success", self.stdout.getvalue()) + + def test_each_live_episode_attempt_resets_policy_before_robot_use(self) -> None: + class StatefulPolicy: + def __init__(self) -> None: + self.state = 99 + self.reset_calls = 0 + + def reset(self) -> None: + self.reset_calls += 1 + self.state = 0 + + policy = StatefulPolicy() + states_at_attempt: list[int] = [] + outcomes = ["operator_failure", "operator_success"] + + class StatefulRunner: + def __init__(self, **kwargs) -> None: + self.policy = kwargs["policy"] + self.robot = kwargs["robot"] + + def run_episode(self, *, attempt: int) -> EpisodeSummary: + states_at_attempt.append(self.policy.state) + self.policy.state += 1 + self.robot.connect() + self.robot.disconnect() + return EpisodeSummary(attempt=attempt, terminal_reason=outcomes.pop(0)) + + def input_fn(prompt: str) -> str: + if "manual reset" in prompt.lower(): + return MANUAL_RESET_PHRASE + return "I HAVE AN E-STOP OPERATOR" if "E-STOP" in prompt else "WORKSPACE IS EMPTY" + + robots: list[FakeRobot] = [] + status = main( + self.rollout_args() + ["--live", "--episode", "--retry-on-failure"], + dependencies=self.dependencies( + input_fn=input_fn, + policy_factory=lambda bundle, device: policy, + robot_factory=lambda **kwargs: robots.append(FakeRobot()) or robots[-1], + runner_factory=lambda **kwargs: StatefulRunner(**kwargs), + ), + ) + + self.assertEqual(status, 0) + self.assertEqual(policy.reset_calls, 2) + self.assertEqual(states_at_attempt, [0, 0]) + self.assertEqual(len(robots), 2) + + def test_second_operator_failure_never_produces_a_third_attempt(self) -> None: + robots: list[FakeRobot] = [] + attempts: list[int] = [] + outcomes = ["operator_failure", "operator_failure"] + + def input_fn(prompt: str) -> str: + if "manual reset" in prompt.lower(): + return MANUAL_RESET_PHRASE + return ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ) + + deps = self.dependencies( + input_fn=input_fn, + robot_factory=lambda **kwargs: robots.append(FakeRobot()) or robots[-1], + runner_factory=lambda **kwargs: ScriptedEpisodeRunner( + outcomes=outcomes, + attempts=attempts, + **kwargs, + ), + ) + + status = main( + self.rollout_args() + + ["--live", "--episode", "--retry-on-failure"], + dependencies=deps, + ) + + self.assertEqual(status, 0) + self.assertEqual(attempts, [1, 2]) + self.assertEqual(len(robots), 2) + + def test_safety_fault_is_never_retried(self) -> None: + robots: list[FakeRobot] = [] + attempts: list[int] = [] + outcomes = ["safety_fault"] + prompts: list[str] = [] + + def input_fn(prompt: str) -> str: + prompts.append(prompt) + return ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ) + + deps = self.dependencies( + input_fn=input_fn, + robot_factory=lambda **kwargs: robots.append(FakeRobot()) or robots[-1], + runner_factory=lambda **kwargs: ScriptedEpisodeRunner( + outcomes=outcomes, + attempts=attempts, + **kwargs, + ), + ) + + status = main( + self.rollout_args() + + ["--live", "--episode", "--retry-on-failure"], + dependencies=deps, + ) + + self.assertEqual(status, 1) + self.assertEqual(attempts, [1]) + self.assertEqual(len(robots), 1) + self.assertFalse(any("manual reset" in prompt.lower() for prompt in prompts)) + + def test_live_rejects_either_incorrect_phrase_before_connect(self) -> None: + cases = ( + ("wrong", "WORKSPACE IS EMPTY"), + ("I HAVE AN E-STOP OPERATOR", "wrong"), + ) + for replies in cases: + with self.subTest(replies=replies): + robot = FakeRobot() + answers = iter(replies) + self.stderr.seek(0) + self.stderr.truncate() + deps = self.dependencies( + input_fn=lambda prompt: next(answers), + robot_factory=lambda **kwargs: robot, + ) + status = main( + self.rollout_args() + ["--live"], dependencies=deps + ) + self.assertEqual(status, 2) + self.assertEqual(robot.connect_count, 0) + + def test_serial_ownership_fails_closed_before_prompt_or_connect(self) -> None: + robot = FakeRobot() + checked: list[str] = [] + deps = self.dependencies( + input_fn=lambda prompt: self.fail("prompted before serial gate"), + robot_factory=lambda **kwargs: robot, + serial_port_is_free=lambda port: checked.append(port) or False, + ) + + status = main(self.rollout_args() + ["--live"], dependencies=deps) + + self.assertEqual(status, 2) + self.assertEqual(checked, ["/dev/fake-follower"]) + self.assertEqual(robot.connect_count, 0) + + def test_dummy_shadow_validates_and_logs_canonical_digest_without_auth_claim(self) -> None: + profile_path = self.write_profile() + robot = FakeRobot() + captured_profiles: list[dict[str, object]] = [] + deps = self.dependencies( + robot_factory=lambda **kwargs: captured_profiles.append( + kwargs["profile_snapshot"] + ) + or robot, + ) + log_path = self.root / "runs" / "policy" / "dummy.jsonl" + + status = main( + [ + "shadow", + "--dummy-hold", + "--profile", + str(profile_path), + "--cycles", + "2", + "--runtime-root", + str(self.root / "runtime"), + "--arm-config", + str(self.root / "arm.yaml"), + "--workspace-config", + str(self.root / "workspace.yaml"), + "--workspace-calibration", + str(self.root / "calibration.json"), + "--log-path", + str(log_path), + ], + dependencies=deps, + ) + + self.assertEqual(status, 0) + self.assertEqual(captured_profiles, [self.profile]) + metadata = json.loads(log_path.read_text().splitlines()[0]) + self.assertEqual(metadata["profile_digest"], canonical_profile_digest(self.profile)) + self.assertEqual(metadata["profile_authentication"], "standalone-untrusted") + self.assertIsNone(metadata["checkpoint"]) + self.assertIsNone(metadata["checkpoint_digest"]) + self.assertNotIn("verified", metadata["profile_authentication"]) + self.assertEqual(robot.send_count, 0) + + def test_invalid_dummy_profile_is_rejected_before_any_rollout_factory(self) -> None: + profile = make_profile() + profile["camera_defaults"]["front"]["width"] = 641 # type: ignore[index] + profile_path = self.write_profile(profile) + called: list[str] = [] + deps = self.dependencies( + robot_factory=lambda **kwargs: called.append("robot"), + guard_factory=lambda **kwargs: called.append("guard"), + safety_factory=lambda *args, **kwargs: called.append("safety"), + ) + + status = main( + [ + "shadow", + "--dummy-hold", + "--profile", + str(profile_path), + ], + dependencies=deps, + ) + + self.assertEqual(status, 2) + self.assertEqual(called, []) + self.assertIn("640x480", self.stderr.getvalue()) + + def test_runner_construction_failure_is_balanced_and_prints_fault_summary(self) -> None: + robot = FakeRobot() + deps = self.dependencies( + input_fn=lambda prompt: ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ), + robot_factory=lambda **kwargs: robot, + runner_factory=lambda **kwargs: (_ for _ in ()).throw( + RuntimeError("runner construction exploded") + ), + ) + + status = main(self.rollout_args() + ["--live"], dependencies=deps) + + self.assertEqual(status, 1) + self.assertEqual(robot.connect_count, 0) + self.assertEqual(robot.disconnect_count, 0) + self.assertEqual(robot.send_count, 0) + summary = self.stdout.getvalue() + self.assertIn("terminal_reason=fault", summary) + self.assertIn("primary_fault=runner construction failed: runner construction exploded", summary) + self.assertIn("cleanup_fault=none", summary) + self.assertIn("audit_fault=none", summary) + + def test_runner_run_exception_prints_summary_without_motion(self) -> None: + robot = FakeRobot() + deps = self.dependencies( + input_fn=lambda prompt: ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ), + robot_factory=lambda **kwargs: robot, + runner_factory=lambda **kwargs: SimpleNamespace( + run=lambda cycles: (_ for _ in ()).throw( + RuntimeError("runner run exploded") + ) + ), + ) + + status = main(self.rollout_args() + ["--live"], dependencies=deps) + + self.assertEqual(status, 1) + self.assertEqual(robot.connect_count, 0) + self.assertEqual(robot.disconnect_count, 0) + self.assertEqual(robot.send_count, 0) + self.assertIn( + "primary_fault=runner execution failed: runner run exploded", + self.stdout.getvalue(), + ) + + def test_preflight_connect_and_cleanup_faults_are_separated_in_summary(self) -> None: + robot = FailingConnectRobot( + cleanup_error=RuntimeError("disconnect exploded") + ) + deps = self.dependencies( + input_fn=lambda prompt: ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ), + robot_factory=lambda **kwargs: robot, + runner_factory=lambda **kwargs: ConnectThenPropagateRunner(**kwargs), + ) + + status = main(self.rollout_args() + ["--live"], dependencies=deps) + + self.assertEqual(status, 1) + self.assertEqual(robot.connect_count, 1) + self.assertEqual(robot.disconnect_count, 1) + self.assertEqual(robot.send_count, 0) + summary = self.stdout.getvalue() + self.assertIn("primary_fault=runner execution failed: connect exploded", summary) + self.assertIn("cleanup_fault=disconnect exploded", summary) + self.assertIn("audit_fault=none", summary) + + def test_real_rollout_processor_artifacts_block_policy_and_hardware(self) -> None: + cases = ("missing_config", "missing_state") + for case in cases: + with self.subTest(case=case): + self.write_processor_configs(self.bundle.path) + if case == "missing_config": + (self.bundle.path / "postprocessor_config.json").unlink() + else: + (self.bundle.path / "preprocessor_config.json").write_text( + json.dumps( + {"steps": [{"state_file": "state/missing.bin"}]} + ), + encoding="utf-8", + ) + calls: list[str] = [] + self.stderr.seek(0) + self.stderr.truncate() + deps = self.dependencies( + input_fn=lambda prompt: self.fail("prompted after processor fault"), + policy_factory=lambda bundle, device: calls.append("policy"), + robot_factory=lambda **kwargs: calls.append("robot"), + ) + + status = main( + self.rollout_args() + ["--live"], dependencies=deps + ) + + self.assertEqual(status, 2) + self.assertEqual(calls, []) + self.assertRegex( + self.stderr.getvalue(), + "postprocessor_config.json|state/missing.bin", + ) + + def test_explicit_log_path_rejects_traversal_protected_roots_and_symlinks(self) -> None: + allowed = self.root / "runs" / "policy" + allowed.mkdir(parents=True) + credential = self.root / "config" / "credentials.env" + credential.parent.mkdir() + credential.write_text("SECRET=unchanged", encoding="utf-8") + symlink = allowed / "linked.jsonl" + symlink.symlink_to(credential) + disallowed = ( + self.root / "data" / "episodes.jsonl", + self.root / "models" / "checkpoint.jsonl", + self.root / "config" / "calibration.jsonl", + allowed / ".." / "escaped.jsonl", + symlink, + allowed / "credentials.jsonl", + allowed / ".env.jsonl", + ) + for path in disallowed: + with self.subTest(path=path): + calls: list[str] = [] + self.stderr.seek(0) + self.stderr.truncate() + args = self.rollout_args() + args[args.index("--log-path") + 1] = str(path) + deps = self.dependencies( + policy_factory=lambda bundle, device: calls.append("policy"), + robot_factory=lambda **kwargs: calls.append("robot"), + ) + + status = main(args + ["--live"], dependencies=deps) + + self.assertEqual(status, 2) + self.assertEqual(calls, []) + self.assertIn("runs/policy", self.stderr.getvalue()) + self.assertEqual( + credential.read_text(encoding="utf-8"), "SECRET=unchanged" + ) + + def test_default_log_path_is_timestamped_under_injected_runs_policy(self) -> None: + args = self.rollout_args() + index = args.index("--log-path") + del args[index : index + 2] + deps = self.dependencies( + input_fn=lambda prompt: ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ) + ) + + status = main(args + ["--live"], dependencies=deps) + + expected = ( + self.root / "runs" / "policy" / "20260718T212223Z.jsonl" + ).resolve() + self.assertEqual(status, 0) + self.assertTrue(expected.is_file()) + self.assertIn(f"jsonl_path={expected}", self.stdout.getvalue()) + + def test_existing_rollout_log_is_refused_without_overwrite(self) -> None: + log_path = self.root / "runs" / "policy" / "rollout.jsonl" + log_path.parent.mkdir(parents=True) + log_path.write_text("existing audit\n", encoding="utf-8") + deps = self.dependencies( + input_fn=lambda prompt: ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ) + ) + + status = main(self.rollout_args() + ["--live"], dependencies=deps) + + self.assertEqual(status, 2) + self.assertEqual(log_path.read_text(encoding="utf-8"), "existing audit\n") + + def test_shared_stop_event_set_after_preflight_prevents_policy_and_send(self) -> None: + from p3_vlm_orchestrator.policy_rollout.runner import RolloutRunner + + keyboard = FakeKeyboardStop() + robot = StopOnSecondObservationRobot(keyboard.event) + policy_calls: list[str] = [] + + class Policy: + def predict(self, observation): + policy_calls.append("predict") + return np.zeros((10, 7)) + + runner_calls: list[dict[str, object]] = [] + + def make_runner(**kwargs): + runner_calls.append(kwargs) + return RolloutRunner(**kwargs) + + deps = self.dependencies( + input_fn=lambda prompt: ( + "I HAVE AN E-STOP OPERATOR" + if "E-STOP" in prompt + else "WORKSPACE IS EMPTY" + ), + policy_factory=lambda bundle, device: Policy(), + robot_factory=lambda **kwargs: robot, + runner_factory=make_runner, + keyboard_stop_factory=lambda: keyboard, + ) + + status = main(self.rollout_args() + ["--live"], dependencies=deps) + + self.assertEqual(status, 0) + self.assertIs(runner_calls[0]["stop_requested"], keyboard.event) + self.assertTrue(keyboard.event.is_set()) + self.assertEqual(policy_calls, []) + self.assertEqual(robot.send_count, 0) + self.assertEqual(robot.connect_count, 1) + self.assertEqual(robot.disconnect_count, 1) + self.assertIn("terminal_reason=stop_requested", self.stdout.getvalue()) + + def test_fresh_process_inspect_and_offline_keep_hardware_modules_unloaded(self) -> None: + checkpoint = self.write_checkpoint() + repo_root = Path(__file__).resolve().parents[2] + commands = ( + (["inspect", "--checkpoint", str(checkpoint)], 0), + ( + [ + "offline", + "--checkpoint", + str(checkpoint), + "--dataset", + str(self.root / "missing-dataset"), + "--episodes", + "1", + ], + 2, + ), + ) + for argv, expected_status in commands: + with self.subTest(command=argv[0]): + script = f""" +import sys +from p3_vlm_orchestrator.policy_rollout.cli import main +status = main({argv!r}) +assert status == {expected_status}, status +forbidden = [] +for name in sys.modules: + lower = name.lower() + if ( + name == 'serial' or name.startswith('serial.') + or name == 'cv2' or name.startswith('cv2.') + or name == 'pinocchio' or name.startswith('pinocchio.') + or name == 'lerobot' or name.startswith('lerobot.') + or name == 'torch' or name.startswith('torch.') + or 'rebot_robot' in lower + or lower.startswith('rebotarm_control_py') + or lower.startswith('p1_arm_motion') + ): + forbidden.append(name) +assert not forbidden, forbidden +""" + completed = subprocess.run( + [sys.executable, "-c", script], + cwd=repo_root, + capture_output=True, + text=True, + check=False, + ) + self.assertEqual( + completed.returncode, + 0, + completed.stdout + completed.stderr, + ) + + +class PreflightRobotAdapterTest(unittest.TestCase): + def test_success_connects_preflights_then_observes_fresh_and_disconnects_once(self) -> None: + first = make_observation() + second = make_observation() + robot = FakeRobot([first, second]) + adapter = PreflightRobotAdapter( + robot=robot, + expected_task=TASK, + profile_snapshot=make_profile(), + ) + + adapter.connect() + self.assertIs(adapter.observe(), second) + adapter.disconnect() + adapter.disconnect() + + self.assertEqual(robot.connect_count, 1) + self.assertEqual(robot.observe_count, 2) + self.assertEqual(robot.disconnect_count, 1) + + def test_each_preflight_mismatch_fails_closed_and_disconnects(self) -> None: + malformed = ( + make_observation(task="different"), + make_observation(state=np.zeros(6)), + make_observation(state=np.array([0, 0, 0, 0, 0, 0, np.nan])), + make_observation(front_shape=(640, 480, 3)), + make_observation(side_shape=(1280, 720, 3)), + make_observation(front_shape=(480, 640)), + make_observation(front_shape=(480, 640, 1)), + make_observation(front_shape=(480, 640, 4)), + make_observation(front_shape=(480, 640, 0)), + make_observation(side_shape=(720, 1280, 1)), + make_observation(side_shape=(720, 1280, 4)), + make_observation(side_shape=(720, 1280, 0)), + ) + for observation in malformed: + with self.subTest( + task=observation.task, + state_shape=np.asarray(observation.state_deg).shape, + front_shape=observation.front.shape, + side_shape=observation.side.shape, + ): + robot = FakeRobot([observation]) + adapter = PreflightRobotAdapter( + robot=robot, + expected_task=TASK, + profile_snapshot=make_profile(), + ) + with self.assertRaises(ValueError): + adapter.connect() + adapter.disconnect() + self.assertEqual(robot.connect_count, 1) + self.assertEqual(robot.disconnect_count, 1) + + +class SerialOwnershipTest(unittest.TestCase): + def test_lsof_argv_and_fail_closed_result_handling(self) -> None: + calls: list[list[str]] = [] + + def result(returncode: int, stdout: str = "", stderr: str = ""): + def run(argv, **kwargs): + calls.append(argv) + return SimpleNamespace( + returncode=returncode, + stdout=stdout, + stderr=stderr, + ) + + return run + + self.assertTrue( + default_serial_port_is_free("/dev/fake", runner=result(1)) + ) + self.assertFalse( + default_serial_port_is_free( + "/dev/fake", runner=result(0, "python 1 user /dev/fake") + ) + ) + self.assertFalse( + default_serial_port_is_free( + "/dev/fake", runner=result(2, stderr="lsof failed") + ) + ) + self.assertFalse( + default_serial_port_is_free( + "/dev/fake", + runner=lambda argv, **kwargs: (_ for _ in ()).throw(OSError("missing")), + ) + ) + self.assertEqual(calls[0], ["lsof", "/dev/fake"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/p3_vlm_orchestrator/tests/test_policy_rollout_runner.py b/p3_vlm_orchestrator/tests/test_policy_rollout_runner.py new file mode 100644 index 0000000..a924910 --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_policy_rollout_runner.py @@ -0,0 +1,798 @@ +from __future__ import annotations + +import json +from pathlib import Path +import tempfile +from threading import Event +import unittest +from unittest.mock import patch + +import numpy as np + +from p3_vlm_orchestrator.policy_rollout.dummy_policy import ( + HoldPositionPolicy, + UnsafePolicy, +) +from p3_vlm_orchestrator.policy_rollout.runner import RolloutRunner +from p3_vlm_orchestrator.policy_rollout.workspace_guard import WorkspaceViolation +from rebot_operator_kit.rollout.contracts import RolloutObservation +from rebot_operator_kit.rollout.safety import SafetyGovernor + + +TASK = "Pick one can and place it in the taped sorting zone" +LIMITS = np.repeat(np.array([[-200.0, 200.0]]), 7, axis=0) +LOG_FIELDS = { + "timestamp_utc", + "monotonic_s", + "event", + "mode", + "cycle", + "task", + "current_state_deg", + "predicted_first_action_deg", + "safety_result", + "send_result_deg", + "inference_latency_s", + "actions_attempted", + "actions_confirmed", + "attempt", + "elapsed_seconds", + "clamp_count", + "terminal_reason", + "fault_reason", + "primary_fault_reason", + "cleanup_fault_reason", + "audit_fault_reason", +} + + +class AdvancingClock: + def __init__(self, start: float = 100.0, step: float = 0.01) -> None: + self.value = start + self.step = step + + def __call__(self) -> float: + value = self.value + self.value += self.step + return value + + +class SequenceClock: + def __init__(self, values: list[float]) -> None: + self.values = iter(values) + self.last = values[-1] + + def __call__(self) -> float: + try: + self.last = next(self.values) + except StopIteration: + self.last += 0.01 + return self.last + + +class FakeRobot: + """Stateful in-memory implementation of the complete robot boundary.""" + + def __init__( + self, + *, + state_deg: np.ndarray | None = None, + captured_monotonic_s: float = 100.0, + stop_after_observation: Event | None = None, + disconnect_error: Exception | None = None, + send_error_after_motion: Exception | None = None, + audit_path: Path | None = None, + ) -> None: + self.state_deg = ( + np.zeros(7, dtype=float) + if state_deg is None + else np.asarray(state_deg, dtype=float).copy() + ) + self.captured_monotonic_s = captured_monotonic_s + self.stop_after_observation = stop_after_observation + self.disconnect_error = disconnect_error + self.send_error_after_motion = send_error_after_motion + self.audit_path = audit_path + self.connected = False + self.connect_count = 0 + self.disconnect_count = 0 + self.observation_states: list[np.ndarray] = [] + self.sent_actions: list[np.ndarray] = [] + self.event_seen_before_send: str | None = None + self.events_seen_before_disconnect: list[str] = [] + + def connect(self) -> None: + if self.connected: + raise RuntimeError("robot is already connected") + self.connected = True + self.connect_count += 1 + + def disconnect(self) -> None: + if self.audit_path is not None and self.audit_path.exists(): + rows = [ + json.loads(line) for line in self.audit_path.read_text().splitlines() + ] + self.events_seen_before_disconnect = [row["event"] for row in rows] + self.connected = False + self.disconnect_count += 1 + if self.disconnect_error is not None: + raise self.disconnect_error + + def observe(self) -> RolloutObservation: + if not self.connected: + raise RuntimeError("robot is disconnected") + state = self.state_deg.copy() + self.observation_states.append(state.copy()) + if self.stop_after_observation is not None: + self.stop_after_observation.set() + return RolloutObservation( + front=np.zeros((2, 2, 3), dtype=np.uint8), + side=np.zeros((2, 2, 3), dtype=np.uint8), + state_deg=state, + task=TASK, + captured_monotonic_s=self.captured_monotonic_s, + ) + + def send_action(self, action_deg: np.ndarray) -> np.ndarray: + if not self.connected: + raise RuntimeError("robot is disconnected") + if self.audit_path is not None: + rows = [ + json.loads(line) for line in self.audit_path.read_text().splitlines() + ] + self.event_seen_before_send = rows[-1]["event"] + actual = np.asarray(action_deg, dtype=float).copy() + self.sent_actions.append(actual.copy()) + self.state_deg = actual.copy() + if self.send_error_after_motion is not None: + raise self.send_error_after_motion + return actual + + +class MalformedStateRobot(FakeRobot): + def observe(self) -> RolloutObservation: + observation = super().observe() + return RolloutObservation( + front=observation.front, + side=observation.side, + state_deg=np.array(["not-a-number"] * 7, dtype=object), + task=observation.task, + captured_monotonic_s=observation.captured_monotonic_s, + ) + + +class FailingAuditFile: + """File-like audit sink that fails one write or flush operation.""" + + def __init__(self, operation: str) -> None: + self.operation = operation + self.failure_pending = True + self.pending: list[str] = [] + self.durable: list[str] = [] + + def __enter__(self) -> FailingAuditFile: + return self + + def __exit__(self, *args: object) -> None: + self.close() + + def write(self, text: str) -> int: + if self.operation == "write" and self.failure_pending: + self.failure_pending = False + raise OSError("simulated audit write failure") + if self.operation == "terminal_write" and self.failure_pending: + row = json.loads(text) + if row.get("event") == "terminal": + self.failure_pending = False + raise OSError("simulated terminal audit write failure") + self.pending.append(text) + return len(text) + + def flush(self) -> None: + if self.operation == "flush" and self.failure_pending: + self.failure_pending = False + self.pending.clear() + raise OSError("simulated audit flush failure") + self.durable.extend(self.pending) + self.pending.clear() + + def close(self) -> None: + self.flush() + + def rows(self) -> list[dict[str, object]]: + return [ + json.loads(line) + for chunk in self.durable + for line in chunk.splitlines() + ] + + +class ArrayPolicy: + def __init__(self, prediction: np.ndarray) -> None: + self.prediction = prediction + self.calls = 0 + + def predict(self, observation: RolloutObservation) -> np.ndarray: + self.calls += 1 + return self.prediction + + +class RecedingHorizonPolicy: + def __init__(self, first_action_deltas: list[float]) -> None: + self.first_action_deltas = first_action_deltas + self.observed_states: list[np.ndarray] = [] + + def predict(self, observation: RolloutObservation) -> np.ndarray: + state = observation.state_deg.copy() + self.observed_states.append(state) + prediction = np.repeat(state[None, :], 10, axis=0) + prediction[0] = state + self.first_action_deltas[len(self.observed_states) - 1] + prediction[1:] = state + 100.0 + return prediction + + +class RaisingPolicy: + def predict(self, observation: RolloutObservation) -> np.ndarray: + raise RuntimeError("inference exploded") + + +class MutatingPolicy: + def predict(self, observation: RolloutObservation) -> np.ndarray: + observation.state_deg[:] = 100.0 + return np.repeat(observation.state_deg[None, :], 10, axis=0) + + +class RecordingActionGuard: + def __init__( + self, + events: list[str] | None = None, + error: Exception | None = None, + ) -> None: + self.events = events + self.error = error + self.actions: list[np.ndarray] = [] + + def validate(self, action_deg: np.ndarray) -> None: + if self.events is not None: + self.events.append("workspace") + self.actions.append(np.asarray(action_deg, dtype=float).copy()) + action_deg[:] = 999.0 + if self.error is not None: + raise self.error + + +class RolloutRunnerTest(unittest.TestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.log_path = Path(temporary.name) / "rollout.jsonl" + + def make_runner( + self, + *, + policy: object, + robot: FakeRobot, + mode: str = "shadow", + stop_requested: Event | None = None, + clock: AdvancingClock | SequenceClock | None = None, + action_guard: object | None = None, + ) -> RolloutRunner: + return RolloutRunner( + policy=policy, + robot=robot, + safety=SafetyGovernor(LIMITS, mode=mode), + mode=mode, + log_path=self.log_path, + monotonic_clock=clock or AdvancingClock(), + stop_requested=stop_requested or Event(), + action_guard=action_guard, + ) + + def read_log(self) -> list[dict[str, object]]: + return [json.loads(line) for line in self.log_path.read_text().splitlines()] + + def assert_full_fallback_schema( + self, + row: dict[str, object], + *, + event: str, + ) -> None: + self.assertEqual(set(row), LOG_FIELDS | {"failed_event"}) + self.assertEqual(row["event"], event) + self.assertTrue(str(row["timestamp_utc"]).endswith("Z")) + for field in ( + "monotonic_s", + "task", + "current_state_deg", + "predicted_first_action_deg", + "safety_result", + "send_result_deg", + "inference_latency_s", + "primary_fault_reason", + "cleanup_fault_reason", + ): + with self.subTest(field=field): + self.assertIsNone(row[field]) + self.assertIsNotNone(row["fault_reason"]) + self.assertIsNotNone(row["audit_fault_reason"]) + + def test_shadow_mode_never_sends_an_action(self) -> None: + robot = FakeRobot() + runner = self.make_runner(policy=HoldPositionPolicy(), robot=robot) + + summary = runner.run(max_cycles=2) + + self.assertEqual(summary.cycles_completed, 2) + self.assertEqual(summary.actions_sent, 0) + self.assertEqual(robot.sent_actions, []) + + def test_public_run_rejects_unbounded_or_nonpositive_cycle_limits(self) -> None: + invalid_limits = (None, True, False, 0, -1, 1.0, "1") + for invalid in invalid_limits: + with self.subTest(max_cycles=invalid): + stop = Event() + stop.set() + robot = FakeRobot() + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + stop_requested=stop, + ) + + with self.assertRaisesRegex(ValueError, "positive integer"): + runner.run(invalid) # type: ignore[arg-type] + + self.assertEqual(robot.connect_count, 0) + self.assertEqual(robot.disconnect_count, 0) + + def test_shadow_invokes_workspace_guard_on_copy_before_safety_and_never_sends(self) -> None: + events: list[str] = [] + guard = RecordingActionGuard(events) + robot = FakeRobot() + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + action_guard=guard, + ) + original_validate = runner.safety.validate + + def safety_validate(*args, **kwargs): + events.append("safety") + return original_validate(*args, **kwargs) + + runner.safety.validate = safety_validate # type: ignore[method-assign] + + summary = runner.run(max_cycles=1) + + self.assertEqual(events, ["workspace", "safety"]) + self.assertEqual(len(guard.actions), 1) + np.testing.assert_array_equal(guard.actions[0], np.zeros(7)) + self.assertEqual(summary.cycles_completed, 1) + self.assertEqual(robot.sent_actions, []) + + def test_workspace_rejection_faults_shadow_and_live_before_safety_or_send(self) -> None: + for mode in ("shadow", "live"): + with self.subTest(mode=mode): + robot = FakeRobot() + guard = RecordingActionGuard( + error=WorkspaceViolation("tip outside calibrated polygon") + ) + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + mode=mode, + action_guard=guard, + ) + safety_calls = 0 + + def unexpected_safety(*args, **kwargs): + nonlocal safety_calls + safety_calls += 1 + raise AssertionError("safety must run after workspace validation") + + runner.safety.validate = unexpected_safety # type: ignore[method-assign] + + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("workspace safety", summary.primary_fault_reason or "") + self.assertEqual(safety_calls, 0) + self.assertEqual(robot.sent_actions, []) + self.assertIn( + "workspace_safety_fault", + [row["event"] for row in self.read_log()], + ) + self.log_path.unlink() + + def test_invalid_prediction_shape_does_not_invoke_workspace_guard(self) -> None: + guard = RecordingActionGuard() + robot = FakeRobot() + runner = self.make_runner( + policy=ArrayPolicy(np.zeros((1, 6))), + robot=robot, + action_guard=guard, + ) + + summary = runner.run(max_cycles=1) + + self.assertIn("prediction shape", summary.primary_fault_reason or "") + self.assertEqual(guard.actions, []) + + def test_hold_position_policy_completes_multiple_fresh_observation_cycles( + self, + ) -> None: + robot = FakeRobot() + runner = self.make_runner(policy=HoldPositionPolicy(), robot=robot) + + summary = runner.run(max_cycles=3) + + self.assertEqual(summary.terminal_reason, "max_cycles") + self.assertIsNone(summary.fault_reason) + self.assertEqual(summary.cycles_completed, 3) + self.assertEqual(len(robot.observation_states), 3) + self.assertEqual(robot.disconnect_count, 1) + + def test_wrong_shape_or_empty_prediction_faults_before_motion(self) -> None: + bad_predictions = ( + np.zeros(7), + np.zeros((1, 6)), + np.zeros((0, 7)), + np.zeros((1, 7, 1)), + ) + + for prediction in bad_predictions: + with self.subTest(shape=prediction.shape): + robot = FakeRobot() + runner = self.make_runner( + policy=ArrayPolicy(prediction), + robot=robot, + ) + + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("prediction shape", summary.fault_reason or "") + self.assertEqual(robot.sent_actions, []) + self.assertEqual(robot.disconnect_count, 1) + + def test_stale_observation_faults_before_motion(self) -> None: + robot = FakeRobot(captured_monotonic_s=0.0) + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + mode="live", + ) + + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("stale", summary.fault_reason or "") + self.assertEqual(robot.sent_actions, []) + + def test_policy_cannot_mutate_current_state_to_bypass_live_safety(self) -> None: + robot = FakeRobot() + runner = self.make_runner( + policy=MutatingPolicy(), + robot=robot, + mode="live", + ) + + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("delta", summary.primary_fault_reason or "") + self.assertEqual(summary.actions_attempted, 0) + self.assertEqual(robot.sent_actions, []) + observation_row = next( + row for row in self.read_log() if row["event"] == "observation" + ) + self.assertEqual(observation_row["current_state_deg"], [0.0] * 7) + + def test_freshness_is_sampled_immediately_before_initial_safety(self) -> None: + clock = SequenceClock([100.0, 100.0, 100.05, 100.30, 100.31]) + robot = FakeRobot() + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + mode="live", + clock=clock, + ) + + summary = runner.run(max_cycles=1) + + self.assertIn("stale", summary.primary_fault_reason or "") + self.assertEqual(summary.actions_attempted, 0) + self.assertEqual(robot.sent_actions, []) + + def test_live_send_boundary_rechecks_observation_freshness(self) -> None: + clock = SequenceClock( + [100.0, 100.0, 100.05, 100.10, 100.20, 100.30, 100.31] + ) + robot = FakeRobot() + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + mode="live", + clock=clock, + ) + + summary = runner.run(max_cycles=1) + + self.assertIn("stale", summary.primary_fault_reason or "") + self.assertEqual(summary.actions_attempted, 0) + self.assertEqual(summary.actions_confirmed, 0) + self.assertEqual(robot.sent_actions, []) + events = [row["event"] for row in self.read_log()] + self.assertIn("send_intent", events) + self.assertIn("send_cancelled", events) + + def test_stop_event_exits_and_disconnects_cleanly(self) -> None: + stop_requested = Event() + stop_requested.set() + robot = FakeRobot() + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + stop_requested=stop_requested, + ) + + summary = runner.run(max_cycles=5) + + self.assertEqual(summary.terminal_reason, "stop_requested") + self.assertIsNone(summary.fault_reason) + self.assertEqual(summary.cycles_completed, 0) + self.assertEqual(robot.observation_states, []) + self.assertEqual(robot.disconnect_count, 1) + + def test_stop_event_after_observation_prevents_inference_and_motion(self) -> None: + stop_requested = Event() + policy = ArrayPolicy(np.zeros((10, 7))) + robot = FakeRobot(stop_after_observation=stop_requested) + runner = self.make_runner( + policy=policy, + robot=robot, + mode="live", + stop_requested=stop_requested, + ) + + summary = runner.run(max_cycles=5) + + self.assertEqual(summary.terminal_reason, "stop_requested") + self.assertEqual(policy.calls, 0) + self.assertEqual(robot.sent_actions, []) + self.assertEqual(robot.disconnect_count, 1) + + def test_live_mode_sends_only_an_accepted_first_action(self) -> None: + policy = RecedingHorizonPolicy([1.0, 2.0]) + robot = FakeRobot() + runner = self.make_runner(policy=policy, robot=robot, mode="live") + + summary = runner.run(max_cycles=2) + + self.assertEqual(summary.actions_sent, 1) + self.assertEqual(summary.actions_attempted, 1) + self.assertEqual(summary.actions_confirmed, 1) + self.assertEqual(summary.cycles_completed, 1) + self.assertIn("safety rejected", summary.fault_reason or "") + self.assertEqual(len(robot.sent_actions), 1) + np.testing.assert_array_equal(robot.sent_actions[0], np.ones(7)) + np.testing.assert_array_equal(policy.observed_states[0], np.zeros(7)) + np.testing.assert_array_equal(policy.observed_states[1], np.ones(7)) + + def test_disconnect_runs_when_inference_raises(self) -> None: + robot = FakeRobot() + runner = self.make_runner(policy=RaisingPolicy(), robot=robot) + + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("inference exploded", summary.fault_reason or "") + self.assertEqual(robot.disconnect_count, 1) + self.assertEqual(robot.sent_actions, []) + fault_row = next(row for row in self.read_log() if row["event"] == "fault") + self.assertEqual(fault_row["event"], "fault") + self.assertIsNotNone(fault_row["inference_latency_s"]) + + def test_send_intent_is_durable_before_adapter_failure_after_motion(self) -> None: + robot = FakeRobot( + send_error_after_motion=RuntimeError("transport acknowledgement lost"), + audit_path=self.log_path, + ) + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + mode="live", + ) + + summary = runner.run(max_cycles=1) + + self.assertEqual(robot.event_seen_before_send, "send_intent") + self.assertEqual(len(robot.sent_actions), 1) + self.assertEqual(summary.actions_attempted, 1) + self.assertEqual(summary.actions_confirmed, 0) + self.assertEqual(summary.actions_sent, 0) + events = [row["event"] for row in self.read_log()] + self.assertLess(events.index("send_intent"), events.index("send_failed")) + self.assertEqual(events[-1], "terminal") + self.assertIn("acknowledgement lost", summary.primary_fault_reason or "") + + def test_jsonl_events_have_full_schema_and_do_not_mutate_arrays(self) -> None: + state = np.zeros(7) + prediction = np.repeat(state[None, :], 10, axis=0) + untouched_state = state.copy() + untouched_prediction = prediction.copy() + robot = FakeRobot(state_deg=state) + runner = self.make_runner( + policy=ArrayPolicy(prediction), + robot=robot, + mode="live", + ) + + runner.run(max_cycles=1) + + rows = self.read_log() + self.assertEqual( + [row["event"] for row in rows], + [ + "observation", + "prediction", + "safety", + "send_intent", + "send_boundary_safety", + "send_confirmed", + "terminal", + ], + ) + for row in rows: + self.assertEqual(set(row), LOG_FIELDS) + self.assertTrue(str(row["timestamp_utc"]).endswith("Z")) + self.assertEqual(row["mode"], "live") + self.assertEqual(rows[1]["predicted_first_action_deg"], [0.0] * 7) + self.assertEqual(rows[2]["safety_result"]["accepted"], True) + self.assertEqual(rows[5]["send_result_deg"], [0.0] * 7) + self.assertEqual(rows[5]["actions_attempted"], 1) + self.assertEqual(rows[5]["actions_confirmed"], 1) + self.assertEqual(rows[-1]["terminal_reason"], "max_cycles") + np.testing.assert_array_equal(state, untouched_state) + np.testing.assert_array_equal(prediction, untouched_prediction) + + def test_unsafe_dummy_policy_is_rejected_without_invalid_json(self) -> None: + robot = FakeRobot() + runner = self.make_runner(policy=UnsafePolicy(), robot=robot, mode="live") + + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("finite", summary.fault_reason or "") + self.assertEqual(robot.sent_actions, []) + prediction_row = next( + row for row in self.read_log() if row["event"] == "prediction" + ) + self.assertEqual(prediction_row["predicted_first_action_deg"], [None] * 7) + + def test_malformed_observation_state_returns_a_structured_fault(self) -> None: + robot = MalformedStateRobot() + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + mode="live", + ) + + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("observation state", summary.primary_fault_reason or "") + self.assertEqual(robot.sent_actions, []) + self.assertEqual(robot.disconnect_count, 1) + rows = self.read_log() + self.assertEqual(rows[-1]["event"], "terminal") + self.assertEqual(rows[-1]["current_state_deg"], None) + + def test_audit_write_or_flush_failure_uses_guarded_fallback(self) -> None: + for operation in ("write", "flush"): + with self.subTest(operation=operation): + audit_file = FailingAuditFile(operation) + robot = FakeRobot() + runner = self.make_runner( + policy=HoldPositionPolicy(), + robot=robot, + mode="live", + ) + + with patch.object(Path, "open", return_value=audit_file): + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + self.assertIn("audit log", summary.audit_fault_reason or "") + self.assertEqual(summary.actions_attempted, 0) + self.assertEqual(robot.sent_actions, []) + self.assertEqual(robot.disconnect_count, 1) + rows = audit_file.rows() + fallback_row = next( + row for row in rows if row["event"] == "fault_fallback" + ) + self.assert_full_fallback_schema( + fallback_row, + event="fault_fallback", + ) + self.assertEqual(rows[-1]["event"], "terminal") + + def test_terminal_fallback_retains_the_full_safe_schema(self) -> None: + audit_file = FailingAuditFile("terminal_write") + robot = FakeRobot() + runner = self.make_runner(policy=HoldPositionPolicy(), robot=robot) + + with patch.object(Path, "open", return_value=audit_file): + summary = runner.run(max_cycles=1) + + self.assertEqual(summary.terminal_reason, "fault") + fallback_row = next( + row + for row in audit_file.rows() + if row["event"] == "terminal_fallback" + ) + self.assert_full_fallback_schema( + fallback_row, + event="terminal_fallback", + ) + self.assertEqual(fallback_row["terminal_reason"], "fault") + + def test_disconnect_fault_is_separate_and_terminal_is_after_cleanup(self) -> None: + robot = FakeRobot( + disconnect_error=RuntimeError("disconnect transport failed"), + audit_path=self.log_path, + ) + runner = self.make_runner(policy=RaisingPolicy(), robot=robot) + + summary = runner.run(max_cycles=1) + + self.assertIn("inference exploded", summary.primary_fault_reason or "") + self.assertIn("disconnect transport", summary.cleanup_fault_reason or "") + self.assertIn("inference exploded", summary.fault_reason or "") + self.assertEqual(robot.disconnect_count, 1) + self.assertNotIn("terminal", robot.events_seen_before_disconnect) + rows = self.read_log() + terminal_rows = [row for row in rows if row["event"] == "terminal"] + self.assertEqual(len(terminal_rows), 1) + self.assertIs(rows[-1], terminal_rows[0]) + self.assertIn( + "disconnect transport", terminal_rows[0]["cleanup_fault_reason"] + ) + + +class DummyPolicyTest(unittest.TestCase): + def test_hold_position_returns_ten_independent_state_copies(self) -> None: + state = np.arange(7, dtype=float) + observation = RolloutObservation( + front=np.zeros((1, 1, 3)), + side=np.zeros((1, 1, 3)), + state_deg=state, + task=TASK, + captured_monotonic_s=1.0, + ) + + prediction = HoldPositionPolicy().predict(observation) + + self.assertEqual(prediction.shape, (10, 7)) + np.testing.assert_array_equal(prediction, np.repeat(state[None, :], 10, axis=0)) + prediction[0, 0] = -999.0 + self.assertEqual(state[0], 0.0) + self.assertEqual(prediction[1, 0], 0.0) + + def test_unsafe_policy_returns_ten_nan_actions(self) -> None: + observation = RolloutObservation( + front=np.zeros((1, 1, 3)), + side=np.zeros((1, 1, 3)), + state_deg=np.zeros(7), + task=TASK, + captured_monotonic_s=1.0, + ) + + prediction = UnsafePolicy().predict(observation) + + self.assertEqual(prediction.shape, (10, 7)) + self.assertTrue(np.all(np.isnan(prediction))) + + +if __name__ == "__main__": + unittest.main() diff --git a/p3_vlm_orchestrator/tests/test_rebot_policy_robot.py b/p3_vlm_orchestrator/tests/test_rebot_policy_robot.py new file mode 100644 index 0000000..7f60750 --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_rebot_policy_robot.py @@ -0,0 +1,719 @@ +from __future__ import annotations + +from hashlib import sha256 +import json +from pathlib import Path +from types import SimpleNamespace +import tempfile +import unittest + +import numpy as np + +from p3_vlm_orchestrator.policy_rollout.rebot_robot import ReBotPolicyRobot + + +JOINTS = ( + "shoulder_pan", + "shoulder_lift", + "elbow_flex", + "wrist_flex", + "wrist_yaw", + "wrist_roll", + "gripper", +) +DIRECTIONS = (-1.0, -1.0, 1.0, 1.0, 1.0, -1.0, -6.0) +LIMITS = ( + (-145.0, 145.0), + (-170.0, 0.0), + (-200.0, 0.0), + (-80.0, 90.0), + (-90.0, 90.0), + (-90.0, 90.0), + (-270.0, 0.0), +) +TASK = "Pick up one can and place it in the taped sorting zone" + + +class FakeBusResource: + def __init__( + self, + *, + disable_error: Exception | None = None, + close_error: Exception | None = None, + ) -> None: + self.disable_error = disable_error + self.close_error = close_error + self.disable_count = 0 + self.close_count = 0 + + def disable_all(self) -> None: + self.disable_count += 1 + if self.disable_error is not None: + raise self.disable_error + + def close(self) -> None: + self.close_count += 1 + if self.close_error is not None: + raise self.close_error + + +class FakeMotorResource: + def __init__( + self, + *, + disable_error: Exception | None = None, + close_error: Exception | None = None, + ) -> None: + self.disable_error = disable_error + self.close_error = close_error + self.disable_count = 0 + self.close_count = 0 + + def disable(self) -> None: + self.disable_count += 1 + if self.disable_error is not None: + raise self.disable_error + + def close(self) -> None: + self.close_count += 1 + if self.close_error is not None: + raise self.close_error + + +class FakeCameraResource: + def __init__( + self, + *, + connected: bool, + disconnect_error: Exception | None = None, + ) -> None: + self.is_connected = connected + self.disconnect_error = disconnect_error + self.disconnect_count = 0 + + def disconnect(self) -> None: + self.disconnect_count += 1 + self.is_connected = False + if self.disconnect_error is not None: + raise self.disconnect_error + + +class FakeFollower: + name = "seeed_b601_dm_follower" + + def __init__(self, config, *, motor_order=JOINTS, camera_keys=("front", "side")): + self.config = config + self.motor_names = list(motor_order) + self.cameras = {key: config.cameras[key] for key in camera_keys} + self.calibration_fpath = config.calibration_dir / f"{config.id}.json" + self.calibration = json.loads(self.calibration_fpath.read_text()) + self.is_connected = False + self.connect_count = 0 + self.disconnect_count = 0 + self.last_command = None + self.observation = { + "front": np.arange(12, dtype=np.uint8).reshape(2, 2, 3), + "side": np.arange(36, dtype=np.uint8).reshape(3, 4, 3), + **{f"{name}.pos": float(index) for index, name in enumerate(JOINTS)}, + } + self.return_override = None + + def connect(self, *, calibrate=True): + self.connect_count += 1 + self.connect_calibrate = calibrate + self.is_connected = True + + def disconnect(self): + self.disconnect_count += 1 + self.is_connected = False + + def get_observation(self): + if not self.is_connected: + raise RuntimeError("disconnected") + return self.observation + + def send_action(self, command): + if not self.is_connected: + raise RuntimeError("disconnected") + self.last_command = dict(command) + if self.return_override is not None: + return self.return_override + return { + f"{name}.pos": command[f"{name}.pos"] + * self.config.joint_directions[name] + for name in JOINTS + } + + +class FakeBackend: + def __init__(self): + self.direction_override = None + self.limit_override = None + self.motor_order = JOINTS + self.camera_keys = ("front", "side") + self.calibration_path_override = None + self.velocity_override = None + self.follower_direction_override = None + self.max_relative_target_override = None + self.disable_torque_override = None + self.loaded_calibration_as_objects = False + self.loaded_calibration_mutator = None + self.follower = None + self.config_kwargs = None + + def make_camera_config(self, **kwargs): + return SimpleNamespace(**kwargs) + + def make_follower_config(self, **kwargs): + self.config_kwargs = kwargs + directions = dict(zip(JOINTS, DIRECTIONS, strict=True)) + limits = dict(zip(JOINTS, LIMITS, strict=True)) + if self.direction_override is not None: + directions = self.direction_override + if self.limit_override is not None: + limits = self.limit_override + config = SimpleNamespace( + **kwargs, + motor_can_ids={name: (index + 1, index + 17) for index, name in enumerate(JOINTS)}, + joint_directions=directions, + joint_limits=limits, + ) + if self.velocity_override is not None: + config.pos_vel_velocity = self.velocity_override + if self.max_relative_target_override is not None: + config.max_relative_target = self.max_relative_target_override + if self.disable_torque_override is not None: + config.disable_torque_on_disconnect = self.disable_torque_override + return config + + def make_follower(self, config): + follower = FakeFollower( + config, + motor_order=self.motor_order, + camera_keys=self.camera_keys, + ) + if self.calibration_path_override is not None: + follower.calibration_fpath = self.calibration_path_override + if self.follower_direction_override is not None: + follower.config = SimpleNamespace(**vars(config)) + follower.config.joint_directions = self.follower_direction_override + if self.loaded_calibration_as_objects: + follower.calibration = { + name: SimpleNamespace(**entry) + for name, entry in follower.calibration.items() + } + if self.loaded_calibration_mutator is not None: + self.loaded_calibration_mutator(follower.calibration) + self.follower = follower + return follower + + +class ReBotPolicyRobotTest(unittest.TestCase): + def setUp(self) -> None: + self.tempdir = tempfile.TemporaryDirectory() + self.runtime_root = Path(self.tempdir.name) + self.backend = FakeBackend() + self.profile = self._make_profile() + self.ports = [SimpleNamespace(device="/dev/fake-follower", vid=0x2E88, pid=0x4603)] + + def tearDown(self) -> None: + self.tempdir.cleanup() + + def _write_runtime_file(self, relative_path: str, payload: bytes) -> str: + path = self.runtime_root / relative_path + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(payload) + return sha256(payload).hexdigest() + + def _make_profile(self) -> dict: + calibration_relative = ( + "lerobot-home/calibration/robots/seeed_b601_dm_follower/follower1.json" + ) + calibration = { + name: { + "id": index + 1, + "drive_mode": 0, + "homing_offset": 0.0, + "range_min": -180.0, + "range_max": 180.0, + } + for index, name in enumerate(JOINTS) + } + calibration_digest = self._write_runtime_file( + calibration_relative, + json.dumps(calibration).encode(), + ) + runtime_entries = {} + for key, relative in ( + ("follower_driver_contract", "driver/config.py"), + ("follower_base_implementation", "driver/base.py"), + ("follower_dm_implementation", "driver/dm.py"), + ): + runtime_entries[key] = { + "runtime_relative_path": relative, + "sha256": self._write_runtime_file(relative, key.encode()), + } + return { + "calibration": { + "follower": { + "type": "seeed_b601_dm_follower", + "id": "follower1", + "runtime_relative_path": calibration_relative, + "sha256": calibration_digest, + }, + **runtime_entries, + }, + "hardware_identity": {"follower_usb": {"vid": 0x2E88, "pid": 0x4603}}, + "coordinate_contract": { + "action_dimension": 7, + "joints": [ + { + "name": name, + "feature": f"{name}.pos", + "leader_to_follower_scale": direction, + "soft_limit_degrees": list(limit), + } + for name, direction, limit in zip(JOINTS, DIRECTIONS, LIMITS, strict=True) + ], + }, + "collection_defaults": { + "task": TASK, + "motor_velocity": 2000.0, + "gripper_force": 0.05, + }, + "camera_defaults": { + "front": { + "recording_key": "observation.images.front", + "index": 0, + "width": 2, + "height": 2, + "fps": 30, + }, + "side": { + "recording_key": "observation.images.side", + "index": 1, + "width": 4, + "height": 3, + "fps": 30, + }, + }, + } + + def _robot(self, **kwargs) -> ReBotPolicyRobot: + return ReBotPolicyRobot.from_profile( + profile_snapshot=self.profile, + runtime_root=self.runtime_root, + speed_scale=0.10, + monotonic_clock=lambda: 123.456, + backend=self.backend, + serial_ports_provider=lambda: list(self.ports), + **kwargs, + ) + + def test_discovers_one_authenticated_follower_and_builds_locked_config_without_connecting( + self, + ) -> None: + robot = self._robot() + + self.assertEqual(self.backend.follower.connect_count, 0) + self.assertFalse(self.backend.follower.is_connected) + self.assertEqual(self.backend.config_kwargs["port"], "/dev/fake-follower") + self.assertEqual(self.backend.config_kwargs["max_relative_target"], 1.5) + self.assertEqual(self.backend.config_kwargs["pos_vel_velocity"], [200.0] * 7) + self.assertTrue(self.backend.config_kwargs["disable_torque_on_disconnect"]) + self.assertEqual(list(self.backend.config_kwargs["cameras"]), ["front", "side"]) + self.assertEqual(robot.task, TASK) + + def test_wrong_no_or_multiple_usb_matches_reject(self) -> None: + cases = ( + [], + [SimpleNamespace(device="/dev/wrong", vid=1, pid=2)], + self.ports + [SimpleNamespace(device="/dev/duplicate", vid=0x2E88, pid=0x4603)], + ) + for ports in cases: + with self.subTest(ports=ports): + self.ports = ports + with self.assertRaisesRegex(ValueError, "exactly one"): + self._robot() + + def test_rejects_speed_outside_first_live_range(self) -> None: + for speed in (0.099, 0.201, np.nan): + with self.subTest(speed=speed), self.assertRaisesRegex(ValueError, "speed_scale"): + ReBotPolicyRobot.from_profile( + profile_snapshot=self.profile, + runtime_root=self.runtime_root, + speed_scale=speed, + backend=self.backend, + serial_ports_provider=lambda: self.ports, + ) + + def test_calibration_or_driver_fingerprint_mismatch_rejects(self) -> None: + entries = ( + "follower", + "follower_driver_contract", + "follower_base_implementation", + "follower_dm_implementation", + ) + for entry in entries: + with self.subTest(entry=entry): + original = self.profile["calibration"][entry]["sha256"] + self.profile["calibration"][entry]["sha256"] = "0" * 64 + try: + with self.assertRaisesRegex(ValueError, "fingerprint"): + self._robot() + finally: + self.profile["calibration"][entry]["sha256"] = original + + def test_rejects_nonintegral_camera_dimensions_without_truncating_profile(self) -> None: + self.profile["camera_defaults"]["front"]["width"] = 2.5 + + with self.assertRaisesRegex(ValueError, "width"): + self._robot() + + def test_rejects_an_extra_camera_mapping(self) -> None: + self.profile["camera_defaults"]["rear"] = { + "recording_key": "observation.images.rear", + "index": 2, + "width": 2, + "height": 2, + "fps": 30, + } + + with self.assertRaisesRegex(ValueError, "exactly front and side"): + self._robot() + + def test_rejects_unknown_scalar_camera_entry(self) -> None: + self.profile["camera_defaults"]["rear"] = 2 + + with self.assertRaisesRegex(ValueError, "extra entries"): + self._robot() + + def test_rejects_duplicate_front_and_side_camera_indices(self) -> None: + self.profile["camera_defaults"]["side"]["index"] = 0 + + with self.assertRaisesRegex(ValueError, "distinct"): + self._robot() + + def test_rejects_boolean_calibration_numbers_even_with_matching_fingerprint(self) -> None: + entry = self.profile["calibration"]["follower"] + path = self.runtime_root / entry["runtime_relative_path"] + calibration = json.loads(path.read_text()) + calibration["shoulder_pan"]["homing_offset"] = True + payload = json.dumps(calibration).encode() + path.write_bytes(payload) + entry["sha256"] = sha256(payload).hexdigest() + + with self.assertRaisesRegex(ValueError, "calibration"): + self._robot() + + def test_rejects_instantiated_velocity_mismatch_before_connect(self) -> None: + self.backend.velocity_override = [199.0] * 7 + + with self.assertRaisesRegex(ValueError, "velocity"): + self._robot() + + self.assertEqual(self.backend.follower.connect_count, 0) + + def test_rejects_authenticated_motor_velocity_other_than_literal_2000(self) -> None: + self.profile["collection_defaults"]["motor_velocity"] = 1999.0 + + with self.assertRaisesRegex(ValueError, "motor velocity.*2000"): + self._robot() + + self.assertIsNone(self.backend.follower) + + def test_rejects_instantiated_relative_target_mismatch_before_connect(self) -> None: + self.backend.max_relative_target_override = 1.6 + + with self.assertRaisesRegex(ValueError, "relative target"): + self._robot() + + self.assertEqual(self.backend.follower.connect_count, 0) + + def test_rejects_instantiated_torque_off_flag_mismatch_before_connect(self) -> None: + self.backend.disable_torque_override = 1 + + with self.assertRaisesRegex(ValueError, "torque"): + self._robot() + + self.assertEqual(self.backend.follower.connect_count, 0) + + def test_rejects_direction_mismatch_on_follower_actual_config(self) -> None: + directions = dict(zip(JOINTS, DIRECTIONS, strict=True)) + directions["wrist_yaw"] = -1.0 + self.backend.follower_direction_override = directions + + with self.assertRaisesRegex(ValueError, "plugin directions"): + self._robot() + + def test_rejects_changed_missing_extra_or_boolean_loaded_calibration_values(self) -> None: + def changed(values): + values["shoulder_pan"]["homing_offset"] = 1.0 + + def missing(values): + values["shoulder_pan"].pop("range_min") + + def extra(values): + values["shoulder_pan"]["unexpected"] = 1.0 + + def boolean(values): + values["shoulder_pan"]["drive_mode"] = False + + for mutate in (changed, missing, extra, boolean): + with self.subTest(mutate=mutate.__name__): + self.backend = FakeBackend() + self.backend.loaded_calibration_mutator = mutate + + with self.assertRaisesRegex(ValueError, "calibration"): + self._robot() + + self.assertEqual(self.backend.follower.connect_count, 0) + + def test_rejects_changed_object_backed_loaded_calibration_value(self) -> None: + def mutate(values): + values["shoulder_pan"].range_max = 181.0 + + self.backend.loaded_calibration_as_objects = True + self.backend.loaded_calibration_mutator = mutate + + with self.assertRaisesRegex(ValueError, "calibration"): + self._robot() + + self.assertEqual(self.backend.follower.connect_count, 0) + + def test_accepts_exact_object_backed_loaded_calibration(self) -> None: + self.backend.loaded_calibration_as_objects = True + + robot = self._robot() + + self.assertEqual(robot.task, TASK) + self.assertEqual(self.backend.follower.connect_count, 0) + + def test_rejects_nonfinite_capture_clock(self) -> None: + robot = ReBotPolicyRobot.from_profile( + profile_snapshot=self.profile, + runtime_root=self.runtime_root, + speed_scale=0.10, + monotonic_clock=lambda: np.nan, + backend=self.backend, + serial_ports_provider=lambda: list(self.ports), + ) + robot.connect() + + with self.assertRaisesRegex(ValueError, "capture"): + robot.observe() + + def test_partial_connect_failure_cleans_bus_motors_and_connected_cameras(self) -> None: + robot = self._robot() + follower = self.backend.follower + bus = FakeBusResource() + motors = { + "shoulder_pan": FakeMotorResource(), + "gripper": FakeMotorResource(), + } + front = FakeCameraResource(connected=True) + side = FakeCameraResource(connected=False) + + def fail_after_front_camera(*, calibrate=True): + follower.bus = bus + follower.motors = motors + follower.cameras = {"front": front, "side": side} + follower.is_connected = False + raise RuntimeError("side camera failed after bus open") + + follower.connect = fail_after_front_camera + + with self.assertRaisesRegex(RuntimeError, "side camera failed") as raised: + robot.connect() + + self.assertIsNone(raised.exception.__cause__) + self.assertEqual(bus.disable_count, 1) + self.assertEqual(bus.close_count, 1) + self.assertIsNone(follower.bus) + for motor in motors.values(): + self.assertEqual(motor.disable_count, 1) + self.assertEqual(motor.close_count, 1) + self.assertEqual(front.disconnect_count, 1) + self.assertEqual(side.disconnect_count, 0) + + def test_partial_connect_preserves_original_error_and_chains_cleanup_failure(self) -> None: + robot = self._robot() + follower = self.backend.follower + bus = FakeBusResource(disable_error=OSError("disable broadcast failed")) + motor = FakeMotorResource(close_error=OSError("motor close failed")) + front = FakeCameraResource( + connected=True, + disconnect_error=OSError("camera disconnect failed"), + ) + + def fail_after_front_camera(*, calibrate=True): + follower.bus = bus + follower.motors = {"shoulder_pan": motor} + follower.cameras = { + "front": front, + "side": FakeCameraResource(connected=False), + } + follower.is_connected = False + raise RuntimeError("original camera open failure") + + follower.connect = fail_after_front_camera + + with self.assertRaisesRegex(RuntimeError, "original camera open failure") as raised: + robot.connect() + + self.assertIsNotNone(raised.exception.__cause__) + self.assertIn("cleanup", str(raised.exception.__cause__).lower()) + self.assertEqual(bus.disable_count, 1) + self.assertEqual(bus.close_count, 1) + self.assertEqual(motor.disable_count, 1) + self.assertEqual(motor.close_count, 1) + self.assertEqual(front.disconnect_count, 1) + + def test_disconnect_cleans_partial_resources_when_aggregate_state_is_false(self) -> None: + robot = self._robot() + robot.connect() + follower = self.backend.follower + bus = FakeBusResource() + motor = FakeMotorResource() + front = FakeCameraResource(connected=True) + follower.bus = bus + follower.motors = {"shoulder_pan": motor} + follower.cameras = { + "front": front, + "side": FakeCameraResource(connected=False), + } + follower.is_connected = False + + robot.disconnect() + + self.assertEqual(bus.disable_count, 1) + self.assertEqual(bus.close_count, 1) + self.assertEqual(motor.disable_count, 1) + self.assertEqual(motor.close_count, 1) + self.assertEqual(front.disconnect_count, 1) + self.assertIsNone(follower.bus) + + def test_plugin_binding_mismatch_rejects_before_connect(self) -> None: + cases = ("direction", "limit", "order", "calibration", "camera") + for case in cases: + with self.subTest(case=case): + self.backend = FakeBackend() + if case == "direction": + self.backend.direction_override = dict(zip(JOINTS, DIRECTIONS, strict=True)) + self.backend.direction_override["wrist_yaw"] = -1.0 + elif case == "limit": + self.backend.limit_override = dict(zip(JOINTS, LIMITS, strict=True)) + self.backend.limit_override["wrist_yaw"] = (-80.0, 80.0) + elif case == "order": + self.backend.motor_order = tuple(reversed(JOINTS)) + elif case == "calibration": + self.backend.calibration_path_override = self.runtime_root / "wrong.json" + else: + self.backend.camera_keys = ("front",) + with self.assertRaisesRegex(ValueError, "plugin|calibration|camera"): + self._robot() + self.assertEqual(self.backend.follower.connect_count, 0) + + def test_observation_copies_locked_images_and_physical_joint_order(self) -> None: + robot = self._robot() + robot.connect() + observation = robot.observe() + + self.assertEqual(self.backend.follower.connect_calibrate, False) + self.assertEqual(observation.task, TASK) + self.assertEqual(observation.captured_monotonic_s, 123.456) + np.testing.assert_array_equal(observation.state_deg, np.arange(7, dtype=float)) + np.testing.assert_array_equal( + observation.front, + self.backend.follower.observation["front"], + ) + np.testing.assert_array_equal(observation.side, self.backend.follower.observation["side"]) + self.assertFalse( + np.shares_memory( + observation.front, + self.backend.follower.observation["front"], + ) + ) + self.assertFalse( + np.shares_memory( + observation.side, + self.backend.follower.observation["side"], + ) + ) + + def test_missing_malformed_or_nonfinite_observation_rejects(self) -> None: + cases = ("front", "side", "joint", "nonfinite") + for case in cases: + with self.subTest(case=case): + self.backend = FakeBackend() + robot = self._robot() + robot.connect() + if case in ("front", "side"): + self.backend.follower.observation.pop(case) + elif case == "joint": + self.backend.follower.observation.pop("wrist_yaw.pos") + else: + self.backend.follower.observation["gripper.pos"] = np.inf + with self.assertRaisesRegex(ValueError, "observation"): + robot.observe() + + def test_inverts_all_seven_physical_actions_before_plugin_direction_mapping(self) -> None: + robot = self._robot() + robot.connect() + requested = np.array((-10.0, -20.0, -30.0, 40.0, 50.0, -60.0, -120.0)) + + returned = robot.send_action(requested) + + expected_command = { + f"{name}.pos": requested[index] / DIRECTIONS[index] + for index, name in enumerate(JOINTS) + } + self.assertEqual(self.backend.follower.last_command, expected_command) + self.assertIn("wrist_yaw.pos", self.backend.follower.last_command) + np.testing.assert_allclose(returned, requested, rtol=0, atol=1e-12) + + def test_zero_direction_or_malformed_nonfinite_disagreeing_plugin_return_rejects(self) -> None: + robot = self._robot() + robot.connect() + requested = np.zeros(7) + feature_keys = [f"{name}.pos" for name in JOINTS] + cases = ( + {key: 0.0 for key in feature_keys[:-1]}, + {key: (np.nan if key == "gripper.pos" else 0.0) for key in feature_keys}, + {key: (1.0 if key == "gripper.pos" else 0.0) for key in feature_keys}, + ) + for value in cases: + with self.subTest(value=value): + self.backend.follower.return_override = value + with self.assertRaisesRegex(ValueError, "returned action"): + robot.send_action(requested) + + self.backend.follower.return_override = None + self.backend.follower.config.joint_directions["gripper"] = 0.0 + with self.assertRaisesRegex(ValueError, "direction"): + robot.send_action(requested) + + def test_mutated_nonfinite_limit_or_overflowing_inversion_rejects_before_send(self) -> None: + for case in ("limit", "direction"): + with self.subTest(case=case): + self.backend = FakeBackend() + robot = self._robot() + robot.connect() + requested = np.zeros(7) + if case == "limit": + self.backend.follower.config.joint_limits["gripper"] = ( + -270.0, + np.nan, + ) + else: + requested[-1] = -1.0 + self.backend.follower.config.joint_directions["gripper"] = 1e-320 + + with self.assertRaisesRegex(ValueError, "limit|invert"): + robot.send_action(requested) + + self.assertIsNone(self.backend.follower.last_command) + + +if __name__ == "__main__": + unittest.main() diff --git a/p3_vlm_orchestrator/tests/test_workspace_guard.py b/p3_vlm_orchestrator/tests/test_workspace_guard.py new file mode 100644 index 0000000..1247b85 --- /dev/null +++ b/p3_vlm_orchestrator/tests/test_workspace_guard.py @@ -0,0 +1,232 @@ +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +import json +from pathlib import Path +import sys +import tempfile +from types import ModuleType +import unittest +from unittest.mock import patch + +import numpy as np + +from p3_vlm_orchestrator.policy_rollout.workspace_guard import ( + CalibratedWorkspaceGuard, + WorkspaceViolation, +) + + +NOW = datetime(2026, 7, 18, 20, 0, tzinfo=timezone.utc) + + +class CalibratedWorkspaceGuardTest(unittest.TestCase): + def setUp(self) -> None: + self.tempdir = tempfile.TemporaryDirectory() + self.root = Path(self.tempdir.name) + self.arm_path = self.root / "arm.yaml" + self.workspace_path = self.root / "workspace.yaml" + self.calibration_path = self.root / "calibration.json" + self._write_files() + + def tearDown(self) -> None: + self.tempdir.cleanup() + + def _write_files( + self, + *, + updated_at: str | None = None, + affine_A: object = ((1.0, 0.0), (0.0, 1.0)), + affine_b: object = (10.0, 20.0), + plane_corners: dict[str, list[float]] | None = None, + width_mm: float = 400.0, + depth_mm: float = 300.0, + transit_height_mm: float = 120.0, + approach_height_mm: float | None = 40.0, + ) -> None: + arm = { + "sdk": {"repo": str(self.root / "sdk")}, + "safety": {"transit_height_mm": transit_height_mm}, + "grasp_heights_mm": {"low": 15.0, "high": 30.0}, + } + if approach_height_mm is not None: + arm["safety"]["approach_height_mm"] = approach_height_mm + self.arm_path.write_text(json.dumps(arm), encoding="utf-8") + self.workspace_path.write_text( + json.dumps({"zone": {"width_mm": width_mm, "depth_mm": depth_mm}}), + encoding="utf-8", + ) + calibration = { + "updated_at": updated_at or (NOW - timedelta(hours=1)).isoformat(), + "aruco": { + "ids": [0, 1, 2, 3], + "plane_mm": plane_corners + or { + "0": [0.0, 0.0], + "1": [400.0, 0.0], + "2": [400.0, 300.0], + "3": [0.0, 300.0], + }, + }, + "plane_to_arm": {"A": affine_A, "b": affine_b}, + } + self.calibration_path.write_text(json.dumps(calibration), encoding="utf-8") + + def _guard(self, fk=lambda action: (210.0, 170.0, 50.0), **kwargs): + return CalibratedWorkspaceGuard.from_files( + arm_config_path=self.arm_path, + workspace_config_path=self.workspace_path, + calibration_path=self.calibration_path, + fk_deg_to_xyz_mm=fk, + current_utc=NOW, + **kwargs, + ) + + def test_accepts_inside_corner_and_reversed_winding_polygon(self) -> None: + self._guard().validate(np.zeros(7)) + self._guard(fk=lambda action: (10.0, 20.0, 5.0)).validate(np.zeros(7)) + + reversed_corners = { + "0": [0.0, 0.0], + "1": [0.0, 300.0], + "2": [400.0, 300.0], + "3": [400.0, 0.0], + } + self._write_files(plane_corners=reversed_corners) + self._guard().validate(np.zeros(7)) + + def test_rejects_tip_outside_xy_or_z_range(self) -> None: + for xyz in ((410.1, 170.0, 50.0), (210.0, 170.0, 4.9), (210.0, 170.0, 160.1)): + with self.subTest(xyz=xyz), self.assertRaises(WorkspaceViolation): + self._guard(fk=lambda action, xyz=xyz: xyz).validate(np.zeros(7)) + + def test_default_z_bounds_use_grasp_transit_and_approach_config(self) -> None: + guard = self._guard() + self.assertEqual(guard.z_min_mm, 5.0) + self.assertEqual(guard.z_max_mm, 160.0) + + self._write_files(approach_height_mm=None) + self.assertEqual(self._guard().z_max_mm, 160.0) + + def test_rejects_invalid_z_range(self) -> None: + self._write_files(transit_height_mm=-50.0, approach_height_mm=40.0) + with self.assertRaisesRegex(ValueError, "Z range"): + self._guard() + + def test_rejects_missing_malformed_or_nonfinite_affine(self) -> None: + for A, b in ((None, None), ([[1, 0]], [0, 0]), ([[1, 0], [0, 1]], [0, np.nan])): + with self.subTest(A=A, b=b): + self._write_files(affine_A=A, affine_b=b) + with self.assertRaisesRegex(ValueError, "affine"): + self._guard() + + def test_rejects_missing_unparseable_expired_or_future_calibration_time(self) -> None: + timestamps = ( + None, + "not-a-time", + (NOW - timedelta(hours=12, seconds=1)).isoformat(), + (NOW + timedelta(minutes=5, seconds=1)).isoformat(), + ) + for timestamp in timestamps: + with self.subTest(timestamp=timestamp): + self._write_files(updated_at=timestamp or "") + if timestamp is None: + data = json.loads(self.calibration_path.read_text()) + data.pop("updated_at") + self.calibration_path.write_text(json.dumps(data)) + with self.assertRaisesRegex(ValueError, "timestamp|stale|future"): + self._guard() + + def test_accepts_configurable_calibration_age_and_callable_clock(self) -> None: + self._write_files(updated_at=(NOW - timedelta(hours=24)).isoformat()) + guard = CalibratedWorkspaceGuard.from_files( + arm_config_path=self.arm_path, + workspace_config_path=self.workspace_path, + calibration_path=self.calibration_path, + fk_deg_to_xyz_mm=lambda action: (210.0, 170.0, 50.0), + current_utc=lambda: NOW, + max_calibration_age_s=25 * 60 * 60, + ) + guard.validate(np.zeros(7)) + + def test_rejects_workspace_dimensions_inconsistent_with_plane_corners(self) -> None: + self._write_files(width_mm=401.0) + with self.assertRaisesRegex(ValueError, "dimensions"): + self._guard() + + def test_default_sdk_fk_is_lazy_converts_first_six_degrees_and_returns_mm(self) -> None: + sdk_repo = self.root / "sdk" + sdk_repo.mkdir() + (sdk_repo / "config").mkdir() + (sdk_repo / "config" / "rebotarm.yaml").write_text( + json.dumps({"hardware_yaml": "rebotarm_dm.yaml"}), + encoding="utf-8", + ) + captured = [] + package = ModuleType("reBotArm_control_py") + package.__path__ = [] # type: ignore[attr-defined] + kinematics = ModuleType("reBotArm_control_py.kinematics") + + def joint_to_pose(q): + captured.append(np.asarray(q).copy()) + return np.array([0.21, 0.17, 0.05]), np.zeros(3) + + kinematics.joint_to_pose = joint_to_pose # type: ignore[attr-defined] + guard = CalibratedWorkspaceGuard.from_files( + arm_config_path=self.arm_path, + workspace_config_path=self.workspace_path, + calibration_path=self.calibration_path, + current_utc=NOW, + ) + self.assertNotIn("reBotArm_control_py.kinematics", sys.modules) + + with patch.dict( + sys.modules, + { + "reBotArm_control_py": package, + "reBotArm_control_py.kinematics": kinematics, + }, + ): + guard.validate(np.array([0.0, 90.0, -90.0, 45.0, -45.0, 180.0, 999.0])) + + np.testing.assert_allclose( + captured[0], + np.radians([0.0, 90.0, -90.0, 45.0, -45.0, 180.0]), + ) + + def test_default_sdk_fk_rejects_missing_or_mismatched_hardware_config(self) -> None: + sdk_repo = self.root / "sdk" + sdk_repo.mkdir() + + with self.assertRaisesRegex(ValueError, "SDK.*hardware"): + CalibratedWorkspaceGuard.from_files( + arm_config_path=self.arm_path, + workspace_config_path=self.workspace_path, + calibration_path=self.calibration_path, + current_utc=NOW, + ) + + (sdk_repo / "config").mkdir() + (sdk_repo / "config" / "rebotarm.yaml").write_text( + json.dumps({"hardware_yaml": "rebotarm_rs.yaml"}), + encoding="utf-8", + ) + with self.assertRaisesRegex(ValueError, "SDK.*hardware"): + CalibratedWorkspaceGuard.from_files( + arm_config_path=self.arm_path, + workspace_config_path=self.workspace_path, + calibration_path=self.calibration_path, + current_utc=NOW, + ) + + def test_rejects_malformed_action_or_fk_return(self) -> None: + with self.assertRaisesRegex(WorkspaceViolation, "shape"): + self._guard().validate(np.zeros(6)) + with self.assertRaisesRegex(WorkspaceViolation, "finite"): + self._guard().validate(np.array([0, 0, 0, 0, 0, 0, np.nan])) + with self.assertRaisesRegex(WorkspaceViolation, "forward kinematics"): + self._guard(fk=lambda action: (1.0, 2.0)).validate(np.zeros(7)) + + +if __name__ == "__main__": + unittest.main() diff --git a/rebot_operator_kit/rollout/__init__.py b/rebot_operator_kit/rollout/__init__.py new file mode 100644 index 0000000..97c9a01 --- /dev/null +++ b/rebot_operator_kit/rollout/__init__.py @@ -0,0 +1 @@ +"""Fail-closed rollout interfaces for ReBot checkpoints and robot control.""" diff --git a/rebot_operator_kit/rollout/checkpoint.py b/rebot_operator_kit/rollout/checkpoint.py new file mode 100644 index 0000000..6b910fe --- /dev/null +++ b/rebot_operator_kit/rollout/checkpoint.py @@ -0,0 +1,176 @@ +"""Dependency-free validation of the Person 3 checkpoint handoff.""" + +from __future__ import annotations + +from dataclasses import dataclass +import hashlib +import json +from pathlib import Path +from typing import Any + + +PROFILE_FILENAME = "rebot_training_profile.json" +REQUIRED_CONFIG_FILENAMES = ( + "config.json", + "preprocessor_config.json", + "postprocessor_config.json", +) +EXPECTED_IMAGE_ORDER = ( + "observation.images.front", + "observation.images.side", +) +EXPECTED_COORDINATE_FRAME = "follower_degrees_after_direction_limits_and_step_cap" +EXPECTED_CONTROL_MODE = "absolute joint pose" +EXPECTED_JOINT_NAMES = ( + "shoulder_pan", + "shoulder_lift", + "elbow_flex", + "wrist_flex", + "wrist_yaw", + "wrist_roll", + "gripper", +) + + +class CheckpointError(ValueError): + """Raised when a checkpoint handoff does not satisfy the rollout contract.""" + + +def _canonical_digest(value: Any) -> str: + payload = json.dumps( + value, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + ).encode("utf-8") + return hashlib.sha256(payload).hexdigest() + + +def _required_object(parent: dict[str, Any], key: str) -> dict[str, Any]: + value = parent.get(key) + if not isinstance(value, dict): + raise CheckpointError(f"Checkpoint profile {key} must be an object") + return value + + +def _positive_integer(value: Any, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise CheckpointError(f"Checkpoint {label} must be a positive integer") + return value + + +@dataclass(frozen=True) +class CheckpointBundle: + path: Path + task: str + action_dimension: int + chunk_size: int + action_steps: int + image_order: tuple[str, str] + profile_digest: str + profile_snapshot: dict[str, Any] + + @classmethod + def load(cls, path: Path) -> CheckpointBundle: + checkpoint = Path(path).expanduser().resolve() + if not checkpoint.is_dir(): + raise CheckpointError(f"Checkpoint directory is missing: {checkpoint}") + + for filename in REQUIRED_CONFIG_FILENAMES: + if not (checkpoint / filename).is_file(): + raise CheckpointError(f"Checkpoint {filename} is missing") + + if not (checkpoint / "model.safetensors").is_file(): + raise CheckpointError("Checkpoint model weights are missing") + + profile_path = checkpoint / PROFILE_FILENAME + if not profile_path.is_file(): + raise CheckpointError(f"Checkpoint {PROFILE_FILENAME} is missing") + try: + sidecar = json.loads(profile_path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + raise CheckpointError( + f"Checkpoint {PROFILE_FILENAME} cannot be read: {exc}" + ) from exc + if not isinstance(sidecar, dict): + raise CheckpointError("Checkpoint profile sidecar must be an object") + + profile = _required_object(sidecar, "profile_snapshot") + declared_profile_digest = sidecar.get("training_profile_digest") + calculated_profile_digest = _canonical_digest(profile) + if declared_profile_digest != calculated_profile_digest: + raise CheckpointError("Checkpoint profile digest does not match its snapshot") + + collection_contract = _required_object(sidecar, "collection_contract") + declared_collection_digest = sidecar.get("collection_contract_digest") + if declared_collection_digest != _canonical_digest(collection_contract): + raise CheckpointError( + "Checkpoint collection contract digest does not match its snapshot" + ) + + collection_defaults = _required_object(profile, "collection_defaults") + locked_task = collection_defaults.get("task") + dataset_task = collection_contract.get("task") + if ( + not isinstance(locked_task, str) + or not locked_task.strip() + or not isinstance(dataset_task, str) + or not dataset_task.strip() + ): + raise CheckpointError("Checkpoint task must be nonempty") + if locked_task != dataset_task: + raise CheckpointError( + "Checkpoint collection task does not match the locked profile task" + ) + + coordinates = _required_object(profile, "coordinate_contract") + training = _required_object(profile, "training_defaults") + if coordinates.get("frame") != EXPECTED_COORDINATE_FRAME: + raise CheckpointError( + f"Checkpoint coordinate frame must be {EXPECTED_COORDINATE_FRAME}" + ) + if coordinates.get("control_mode") != EXPECTED_CONTROL_MODE: + raise CheckpointError( + f"Checkpoint control mode must be {EXPECTED_CONTROL_MODE}" + ) + joints = coordinates.get("joints") + joint_names = ( + [joint.get("name") for joint in joints] + if isinstance(joints, list) + and all(isinstance(joint, dict) for joint in joints) + else None + ) + if joint_names != list(EXPECTED_JOINT_NAMES): + raise CheckpointError( + "Checkpoint joint names must match the seven-joint profile order" + ) + coordinate_dimension = coordinates.get("action_dimension") + training_dimension = training.get("action_dimension") + if ( + isinstance(coordinate_dimension, bool) + or not isinstance(coordinate_dimension, int) + or coordinate_dimension != 7 + or isinstance(training_dimension, bool) + or not isinstance(training_dimension, int) + or training_dimension != 7 + ): + raise CheckpointError("Checkpoint action dimension must be 7") + + chunk_size = _positive_integer(training.get("chunk_size"), "chunk size") + action_steps = _positive_integer( + training.get("n_action_steps"), "action steps" + ) + image_order = training.get("image_order") + if image_order != list(EXPECTED_IMAGE_ORDER): + raise CheckpointError("Checkpoint image order must be front then side") + + return cls( + path=checkpoint, + task=locked_task, + action_dimension=training_dimension, + chunk_size=chunk_size, + action_steps=action_steps, + image_order=EXPECTED_IMAGE_ORDER, + profile_digest=calculated_profile_digest, + profile_snapshot=profile, + ) diff --git a/rebot_operator_kit/rollout/contracts.py b/rebot_operator_kit/rollout/contracts.py new file mode 100644 index 0000000..794bef4 --- /dev/null +++ b/rebot_operator_kit/rollout/contracts.py @@ -0,0 +1,31 @@ +"""Pure interfaces shared by rollout policies, robots, and runners.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Protocol + +import numpy as np + + +@dataclass(frozen=True) +class RolloutObservation: + front: np.ndarray + side: np.ndarray + state_deg: np.ndarray + task: str + captured_monotonic_s: float + + +class PolicyAdapter(Protocol): + def predict(self, observation: RolloutObservation) -> np.ndarray: ... + + +class RobotAdapter(Protocol): + def connect(self) -> None: ... + + def disconnect(self) -> None: ... + + def observe(self) -> RolloutObservation: ... + + def send_action(self, action_deg: np.ndarray) -> np.ndarray: ... diff --git a/rebot_operator_kit/rollout/safety.py b/rebot_operator_kit/rollout/safety.py new file mode 100644 index 0000000..60f5afd --- /dev/null +++ b/rebot_operator_kit/rollout/safety.py @@ -0,0 +1,189 @@ +"""Fail-closed validation for rollout actions before robot execution.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Any + +import numpy as np + + +JOINT_COUNT = 7 +DEFAULT_MAX_OBSERVATION_AGE_S = 0.250 +DEFAULT_MAX_ACTION_DELTA_DEG = 1.5 +DEFAULT_MAX_CONSECUTIVE_INTERVENTIONS = 3 + + +@dataclass(frozen=True) +class SafetyDecision: + accepted: bool + action_deg: np.ndarray | None + reason: str + clamped: bool = False + + +class SafetyFault(RuntimeError): + """Raised once repeated safety interventions latch the governor.""" + + +class SafetyGovernor: + """Validate seven-joint actions against an authenticated profile.""" + + def __init__( + self, + hard_limits_deg: Sequence[Sequence[float]] | np.ndarray, + *, + mode: str, + max_observation_age_s: float = DEFAULT_MAX_OBSERVATION_AGE_S, + max_action_delta_deg: float = DEFAULT_MAX_ACTION_DELTA_DEG, + max_consecutive_interventions: int = DEFAULT_MAX_CONSECUTIVE_INTERVENTIONS, + ) -> None: + if mode not in ("shadow", "live"): + raise ValueError("Safety governor mode must be exactly 'shadow' or 'live'") + + try: + limits = np.asarray(hard_limits_deg, dtype=float).copy() + except (TypeError, ValueError) as exc: + raise ValueError("Safety hard limits must be numeric") from exc + if limits.shape != (JOINT_COUNT, 2): + raise ValueError("Safety hard limits must have shape (7, 2)") + if not np.all(np.isfinite(limits)): + raise ValueError("Safety hard limits must be finite") + if np.any(limits[:, 0] >= limits[:, 1]): + raise ValueError("Each safety hard-limit minimum must be below its maximum") + + if not np.isfinite(max_observation_age_s) or max_observation_age_s < 0: + raise ValueError("Maximum observation age must be finite and nonnegative") + if not np.isfinite(max_action_delta_deg) or max_action_delta_deg <= 0: + raise ValueError("Maximum action delta must be finite and positive") + if ( + isinstance(max_consecutive_interventions, bool) + or not isinstance(max_consecutive_interventions, int) + or max_consecutive_interventions <= 0 + ): + raise ValueError("Maximum consecutive interventions must be positive") + + self.mode = mode + self.hard_limits_deg = limits + self.max_observation_age_s = float(max_observation_age_s) + self.max_action_delta_deg = float(max_action_delta_deg) + self.max_consecutive_interventions = max_consecutive_interventions + self._consecutive_interventions = 0 + self._fault: SafetyFault | None = None + + @classmethod + def from_profile( + cls, + profile_snapshot: Mapping[str, Any], + *, + mode: str, + max_observation_age_s: float = DEFAULT_MAX_OBSERVATION_AGE_S, + max_action_delta_deg: float = DEFAULT_MAX_ACTION_DELTA_DEG, + max_consecutive_interventions: int = DEFAULT_MAX_CONSECUTIVE_INTERVENTIONS, + ) -> SafetyGovernor: + """Build a governor from the checkpoint's authenticated profile snapshot.""" + + try: + coordinate_contract = profile_snapshot["coordinate_contract"] + joints = coordinate_contract["joints"] + limits = [joint["soft_limit_degrees"] for joint in joints] + except (KeyError, TypeError) as exc: + raise ValueError( + "Profile must contain coordinate_contract joint soft limits" + ) from exc + + return cls( + limits, + mode=mode, + max_observation_age_s=max_observation_age_s, + max_action_delta_deg=max_action_delta_deg, + max_consecutive_interventions=max_consecutive_interventions, + ) + + def validate( + self, + current: np.ndarray, + proposed: np.ndarray, + now: float, + observed_at: float, + ) -> SafetyDecision: + """Return an executable copy or reject and count a safety intervention.""" + + if self._fault is not None: + raise self._fault + + try: + current_deg = np.asarray(current, dtype=float).copy() + proposed_deg = np.asarray(proposed, dtype=float).copy() + except (TypeError, ValueError): + return self._reject("current and proposed joint values must be finite numbers") + + if current_deg.shape != (JOINT_COUNT,) or proposed_deg.shape != (JOINT_COUNT,): + return self._reject("current and proposed actions must each have shape (7,)") + if not np.all(np.isfinite(current_deg)) or not np.all(np.isfinite(proposed_deg)): + return self._reject("current and proposed joint values must be finite") + + try: + checked_now = float(now) + checked_observed_at = float(observed_at) + except (TypeError, ValueError): + return self._reject("observation timestamps must be finite") + if not np.isfinite(checked_now) or not np.isfinite(checked_observed_at): + return self._reject("observation timestamps must be finite") + + observation_age = checked_now - checked_observed_at + if observation_age < 0: + return self._reject("observation timestamp is in the future") + if observation_age > self.max_observation_age_s: + return self._reject("observation is stale") + + lower = self.hard_limits_deg[:, 0] + upper = self.hard_limits_deg[:, 1] + if ( + np.any(current_deg < lower) + or np.any(current_deg > upper) + or np.any(proposed_deg < lower) + or np.any(proposed_deg > upper) + ): + return self._reject("current or proposed action is outside hard limits") + + delta = proposed_deg - current_deg + if np.any(np.abs(delta) > self.max_action_delta_deg): + if self.mode == "live": + return self._reject("action delta exceeds the live step cap") + clamped = current_deg + np.clip( + delta, + -self.max_action_delta_deg, + self.max_action_delta_deg, + ) + return self._intervene( + SafetyDecision( + accepted=True, + action_deg=clamped, + reason="action delta clamped to the live step cap", + clamped=True, + ) + ) + + self._consecutive_interventions = 0 + return SafetyDecision( + accepted=True, + action_deg=proposed_deg, + reason="accepted", + ) + + def _reject(self, reason: str) -> SafetyDecision: + return self._intervene( + SafetyDecision(accepted=False, action_deg=None, reason=reason) + ) + + def _intervene(self, decision: SafetyDecision) -> SafetyDecision: + self._consecutive_interventions += 1 + if self._consecutive_interventions >= self.max_consecutive_interventions: + self._fault = SafetyFault( + "Safety governor latched after " + f"{self._consecutive_interventions} consecutive interventions" + ) + raise self._fault + return decision diff --git a/rebot_operator_kit/tests/test_rollout_contract.py b/rebot_operator_kit/tests/test_rollout_contract.py new file mode 100644 index 0000000..c8bdcaf --- /dev/null +++ b/rebot_operator_kit/tests/test_rollout_contract.py @@ -0,0 +1,256 @@ +from __future__ import annotations + +import hashlib +import json +from dataclasses import FrozenInstanceError +from pathlib import Path +import tempfile +import unittest + +import numpy as np + +from rebot_operator_kit.rollout.checkpoint import CheckpointBundle, CheckpointError +from rebot_operator_kit.rollout.contracts import ( + PolicyAdapter, + RobotAdapter, + RolloutObservation, +) + + +FRONT_IMAGE_KEY = "observation.images.front" +SIDE_IMAGE_KEY = "observation.images.side" +TASK = "Pick up the crumpled paper ball and place it in the trash bin" +COORDINATE_FRAME = "follower_degrees_after_direction_limits_and_step_cap" +CONTROL_MODE = "absolute joint pose" +JOINT_NAMES = [ + "shoulder_pan", + "shoulder_lift", + "elbow_flex", + "wrist_flex", + "wrist_yaw", + "wrist_roll", + "gripper", +] + + +def canonical_digest(value: object) -> str: + payload = json.dumps( + value, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + ).encode("utf-8") + return hashlib.sha256(payload).hexdigest() + + +class RolloutObservationTest(unittest.TestCase): + def test_observation_is_immutable(self) -> None: + observation = RolloutObservation( + front=np.zeros((2, 2, 3)), + side=np.zeros((2, 2, 3)), + state_deg=np.zeros(7), + task=TASK, + captured_monotonic_s=123.0, + ) + + with self.assertRaises(FrozenInstanceError): + observation.task = "changed" + + def test_adapter_protocols_expose_the_rollout_boundary(self) -> None: + self.assertTrue(hasattr(PolicyAdapter, "predict")) + for method in ("connect", "disconnect", "observe", "send_action"): + with self.subTest(method=method): + self.assertTrue(hasattr(RobotAdapter, method)) + + +class CheckpointBundleTest(unittest.TestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.checkpoint = Path(temporary.name) + self.profile = { + "schema_version": 1, + "profile_id": "rebot-test-profile", + "profile_version": 1, + "collection_defaults": {"task": TASK}, + "coordinate_contract": { + "frame": COORDINATE_FRAME, + "control_mode": CONTROL_MODE, + "action_dimension": 7, + "joints": [{"name": name} for name in JOINT_NAMES], + }, + "training_defaults": { + "action_dimension": 7, + "chunk_size": 10, + "n_action_steps": 10, + "image_order": [FRONT_IMAGE_KEY, SIDE_IMAGE_KEY], + }, + } + self.collection_contract = {"task": TASK} + self._write_required_files() + self._write_profile_sidecar() + + def _write_json(self, name: str, payload: object) -> None: + (self.checkpoint / name).write_text(json.dumps(payload)) + + def _write_required_files(self) -> None: + self._write_json("config.json", {"policy_type": "molmoact2"}) + self._write_json("preprocessor_config.json", {}) + self._write_json("postprocessor_config.json", {}) + (self.checkpoint / "model.safetensors").write_bytes(b"test weights") + + def _write_profile_sidecar(self, *, digest: str | None = None) -> None: + self._write_json( + "rebot_training_profile.json", + { + "schema_version": 1, + "training_profile_id": self.profile["profile_id"], + "training_profile_version": self.profile["profile_version"], + "training_profile_digest": digest or canonical_digest(self.profile), + "profile_snapshot": self.profile, + "collection_contract": self.collection_contract, + "collection_contract_digest": canonical_digest(self.collection_contract), + }, + ) + + def assert_rejected(self, expected_message: str) -> None: + with self.assertRaisesRegex(CheckpointError, expected_message): + CheckpointBundle.load(self.checkpoint) + + def test_accepts_complete_checkpoint_manifest(self) -> None: + bundle = CheckpointBundle.load(self.checkpoint) + + self.assertEqual(bundle.path, self.checkpoint.resolve()) + self.assertEqual(bundle.task, TASK) + self.assertEqual(bundle.action_dimension, 7) + self.assertEqual(bundle.chunk_size, 10) + self.assertEqual(bundle.action_steps, 10) + self.assertEqual(bundle.image_order, (FRONT_IMAGE_KEY, SIDE_IMAGE_KEY)) + self.assertEqual(bundle.profile_digest, canonical_digest(self.profile)) + self.assertEqual(bundle.profile_snapshot, self.profile) + + def test_rejects_missing_profile(self) -> None: + (self.checkpoint / "rebot_training_profile.json").unlink() + + self.assert_rejected("rebot_training_profile.json.*missing") + + def test_rejects_profile_digest_mismatch(self) -> None: + self._write_profile_sidecar(digest="0" * 64) + + self.assert_rejected("digest.*match") + + def test_rejects_collection_contract_digest_mismatch(self) -> None: + self.collection_contract["task"] = TASK + " safely" + sidecar = json.loads( + (self.checkpoint / "rebot_training_profile.json").read_text() + ) + sidecar["collection_contract"] = self.collection_contract + self._write_json("rebot_training_profile.json", sidecar) + + self.assert_rejected("collection contract digest.*match") + + def test_rejects_task_mismatch(self) -> None: + self.collection_contract["task"] = "Perform a different task" + self._write_profile_sidecar() + + self.assert_rejected("task.*match") + + def test_rejects_blank_task(self) -> None: + self.profile["collection_defaults"]["task"] = " " + self.collection_contract["task"] = " " + self._write_profile_sidecar() + + self.assert_rejected("task.*nonempty") + + def test_rejects_action_dimension_other_than_seven(self) -> None: + self.profile["training_defaults"]["action_dimension"] = 6 + self.profile["coordinate_contract"]["action_dimension"] = 6 + self._write_profile_sidecar() + + self.assert_rejected("action dimension.*7") + + def test_rejects_profile_action_dimension_disagreement(self) -> None: + self.profile["coordinate_contract"]["action_dimension"] = 6 + self._write_profile_sidecar() + + self.assert_rejected("action dimension.*7") + + def test_rejects_reordered_joint_names(self) -> None: + joints = self.profile["coordinate_contract"]["joints"] + joints[0], joints[1] = joints[1], joints[0] + self._write_profile_sidecar() + + self.assert_rejected("joint names.*order") + + def test_rejects_missing_coordinate_frame(self) -> None: + self.profile["coordinate_contract"].pop("frame") + self._write_profile_sidecar() + + self.assert_rejected("coordinate frame") + + def test_rejects_wrong_coordinate_frame(self) -> None: + self.profile["coordinate_contract"]["frame"] = "leader_degrees" + self._write_profile_sidecar() + + self.assert_rejected("coordinate frame") + + def test_rejects_missing_control_mode(self) -> None: + self.profile["coordinate_contract"].pop("control_mode") + self._write_profile_sidecar() + + self.assert_rejected("control mode") + + def test_rejects_wrong_control_mode(self) -> None: + self.profile["coordinate_contract"]["control_mode"] = "velocity" + self._write_profile_sidecar() + + self.assert_rejected("control mode") + + def test_rejects_nonpositive_chunk_size(self) -> None: + self.profile["training_defaults"]["chunk_size"] = 0 + self._write_profile_sidecar() + + self.assert_rejected("chunk size.*positive") + + def test_rejects_nonpositive_action_steps(self) -> None: + self.profile["training_defaults"]["n_action_steps"] = 0 + self._write_profile_sidecar() + + self.assert_rejected("action steps.*positive") + + def test_rejects_wrong_image_order(self) -> None: + self.profile["training_defaults"]["image_order"] = [ + SIDE_IMAGE_KEY, + FRONT_IMAGE_KEY, + ] + self._write_profile_sidecar() + + self.assert_rejected("image order.*front then side") + + def test_rejects_missing_processor_config(self) -> None: + for filename in ("preprocessor_config.json", "postprocessor_config.json"): + with self.subTest(filename=filename): + path = self.checkpoint / filename + path.unlink() + self.assert_rejected(f"{filename}.*missing") + self._write_json(filename, {}) + + def test_rejects_missing_policy_config(self) -> None: + (self.checkpoint / "config.json").unlink() + + self.assert_rejected("config.json.*missing") + + def test_rejects_missing_model_weights(self) -> None: + (self.checkpoint / "model.safetensors").unlink() + + self.assert_rejected("model weights.*missing") + + def test_rejects_unrecognized_binary_as_model_weights(self) -> None: + (self.checkpoint / "model.safetensors").unlink() + (self.checkpoint / "optimizer.bin").write_bytes(b"not policy weights") + + self.assert_rejected("model weights.*missing") + + +if __name__ == "__main__": + unittest.main() diff --git a/rebot_operator_kit/tests/test_rollout_safety.py b/rebot_operator_kit/tests/test_rollout_safety.py new file mode 100644 index 0000000..dd47e57 --- /dev/null +++ b/rebot_operator_kit/tests/test_rollout_safety.py @@ -0,0 +1,204 @@ +from __future__ import annotations + +from dataclasses import FrozenInstanceError +import unittest + +import numpy as np + +from rebot_operator_kit.rollout.safety import ( + SafetyDecision, + SafetyFault, + SafetyGovernor, +) + + +LIMITS = np.array( + [ + [-145.0, 145.0], + [-170.0, 0.0], + [-200.0, 0.0], + [-80.0, 90.0], + [-90.0, 90.0], + [-90.0, 90.0], + [-270.0, 0.0], + ] +) +IN_RANGE = np.array([0.0, -80.0, -100.0, 0.0, 0.0, 0.0, -100.0]) + + +def profile_snapshot() -> dict[str, object]: + return { + "coordinate_contract": { + "joints": [ + {"soft_limit_degrees": bounds.tolist()} for bounds in LIMITS + ] + } + } + + +class SafetyGovernorTest(unittest.TestCase): + def setUp(self) -> None: + self.shadow = SafetyGovernor.from_profile(profile_snapshot(), mode="shadow") + self.live = SafetyGovernor.from_profile(profile_snapshot(), mode="live") + + def validate( + self, + governor: SafetyGovernor, + current: np.ndarray = IN_RANGE, + proposed: np.ndarray = IN_RANGE, + *, + now: float = 10.0, + observed_at: float = 10.0, + ) -> SafetyDecision: + return governor.validate(current, proposed, now, observed_at) + + def assert_rejected_without_replacement(self, decision: SafetyDecision) -> None: + self.assertFalse(decision.accepted) + self.assertIsNone(decision.action_deg) + self.assertFalse(decision.clamped) + + def test_rejects_current_or_proposed_shape_other_than_seven_vector(self) -> None: + bad_shapes = (np.zeros(6), np.zeros((1, 7)), np.zeros(8)) + + for field in ("current", "proposed"): + for bad_value in bad_shapes: + with self.subTest(field=field, shape=bad_value.shape): + governor = SafetyGovernor(LIMITS, mode="live") + values = {"current": IN_RANGE, "proposed": IN_RANGE} + values[field] = bad_value + + decision = self.validate(governor, **values) + + self.assert_rejected_without_replacement(decision) + self.assertIn("shape (7,)", decision.reason) + + def test_rejects_nan_or_infinity_without_replacement(self) -> None: + for field in ("current", "proposed"): + for value in (np.nan, np.inf, -np.inf): + with self.subTest(field=field, value=value): + governor = SafetyGovernor(LIMITS, mode="shadow") + values = { + "current": IN_RANGE.copy(), + "proposed": IN_RANGE.copy(), + } + values[field][0] = value + + decision = self.validate(governor, **values) + + self.assert_rejected_without_replacement(decision) + self.assertIn("finite", decision.reason) + + def test_rejects_observation_older_than_250_milliseconds(self) -> None: + decision = self.validate(self.live, now=10.251, observed_at=10.0) + + self.assert_rejected_without_replacement(decision) + self.assertIn("stale", decision.reason) + + def test_accepts_observation_at_250_millisecond_boundary(self) -> None: + decision = self.validate(self.live, now=10.250, observed_at=10.0) + + self.assertTrue(decision.accepted) + + def test_rejects_current_or_proposed_value_outside_hard_limits(self) -> None: + for field in ("current", "proposed"): + with self.subTest(field=field): + governor = SafetyGovernor(LIMITS, mode="shadow") + values = { + "current": IN_RANGE.copy(), + "proposed": IN_RANGE.copy(), + } + values[field][0] = LIMITS[0, 1] + 0.1 + + decision = self.validate(governor, **values) + + self.assert_rejected_without_replacement(decision) + self.assertIn("hard limits", decision.reason) + + def test_shadow_mode_clamps_excessive_delta_and_accepts_copy(self) -> None: + proposed = IN_RANGE.copy() + proposed[[0, 3]] += np.array([3.0, -2.0]) + untouched_proposed = proposed.copy() + + decision = self.validate(self.shadow, proposed=proposed) + + self.assertTrue(decision.accepted) + self.assertTrue(decision.clamped) + np.testing.assert_allclose( + decision.action_deg, + IN_RANGE + np.array([1.5, 0.0, 0.0, -1.5, 0.0, 0.0, 0.0]), + ) + np.testing.assert_array_equal(proposed, untouched_proposed) + self.assertIsNot(decision.action_deg, proposed) + + def test_live_mode_rejects_excessive_delta_without_replacement(self) -> None: + proposed = IN_RANGE.copy() + proposed[0] += 1.5001 + + decision = self.validate(self.live, proposed=proposed) + + self.assert_rejected_without_replacement(decision) + self.assertIn("delta", decision.reason) + + def test_valid_in_range_action_is_accepted_as_copy(self) -> None: + proposed = IN_RANGE + np.array([1.5, -1.0, 0.5, 0.0, 0.0, 0.0, 0.0]) + + decision = self.validate(self.live, proposed=proposed) + + self.assertTrue(decision.accepted) + self.assertFalse(decision.clamped) + self.assertEqual(decision.reason, "accepted") + np.testing.assert_array_equal(decision.action_deg, proposed) + self.assertIsNot(decision.action_deg, proposed) + + def test_profile_soft_limits_are_enforced_as_hard_rollout_bounds(self) -> None: + proposed = IN_RANGE.copy() + proposed[6] = LIMITS[6, 0] - 0.1 + + decision = self.validate(self.live, proposed=proposed) + + self.assert_rejected_without_replacement(decision) + self.assertIn("hard limits", decision.reason) + + def test_third_consecutive_rejection_latches_fault(self) -> None: + stale = {"now": 10.251, "observed_at": 10.0} + for _ in range(2): + self.assertFalse(self.validate(self.live, **stale).accepted) + + with self.assertRaises(SafetyFault): + self.validate(self.live, **stale) + with self.assertRaises(SafetyFault): + self.validate(self.live) + + def test_third_consecutive_shadow_clamp_latches_fault(self) -> None: + proposed = IN_RANGE.copy() + proposed[0] += 2.0 + for _ in range(2): + self.assertTrue(self.validate(self.shadow, proposed=proposed).clamped) + + with self.assertRaises(SafetyFault): + self.validate(self.shadow, proposed=proposed) + + def test_normal_acceptance_resets_consecutive_interventions(self) -> None: + stale = {"now": 10.251, "observed_at": 10.0} + for _ in range(2): + self.assertFalse(self.validate(self.live, **stale).accepted) + + self.assertTrue(self.validate(self.live).accepted) + + for _ in range(2): + self.assertFalse(self.validate(self.live, **stale).accepted) + self.assertTrue(self.validate(self.live).accepted) + + def test_rejects_unsupported_execution_mode(self) -> None: + with self.assertRaisesRegex(ValueError, "shadow.*live"): + SafetyGovernor(LIMITS, mode="offline") + + def test_safety_decision_is_frozen(self) -> None: + decision = self.validate(self.live) + + with self.assertRaises(FrozenInstanceError): + decision.accepted = False + + +if __name__ == "__main__": + unittest.main()