#!/usr/bin/env python3 """Pure control protocol for the independent TG3 data recorder. The ROS node is deliberately kept as a thin transport adapter. This module owns JSON validation, idempotence, session/capture matching, acknowledgement state, and the bridge-heartbeat watchdog, and can therefore be tested without ROS installed. """ from __future__ import annotations import json import math import time from collections import OrderedDict from collections.abc import Callable, Mapping, Sequence from pathlib import Path from typing import Any, Protocol PROTOCOL_VERSION = 1 ACTIVE_MANAGER_STATES = { "start_pending", "starting", "recording", "stop_pending", "finalizing", } class RecorderBackend(Protocol): """The non-blocking subset exposed by ``DataRecorderManager``.""" def request_start( self, session_id: str, reason: str = "button", *, episode_id: str | None = None, ) -> str | None: ... def request_stop(self, reason: str = "operator_button") -> bool: ... def status(self) -> dict[str, Any]: ... def shutdown( self, *, wait: bool = True, timeout: float | None = None ) -> bool: ... class ProtocolError(ValueError): """A control message is malformed or invalid for the current session.""" def __init__(self, code: str, message: str) -> None: super().__init__(message) self.code = code def _nonempty_text(value: Any, name: str, *, maximum: int = 128) -> str: if ( not isinstance(value, str) or not value or value.strip() != value or len(value) > maximum or any(ord(character) < 0x20 for character in value) ): raise ProtocolError( f"invalid_{name}", f"{name} must be non-empty text up to {maximum} characters", ) return value def _capture_id(value: Any) -> str: capture_id = _nonempty_text(value, "capture_id") if ( not capture_id[0].isalnum() or Path(capture_id).name != capture_id or capture_id in (".", "..") or any( character not in "-_.0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" for character in capture_id ) ): raise ProtocolError( "invalid_capture_id", "capture_id must be one safe path component beginning with a letter or digit", ) return capture_id def _event_sequence(value: Any) -> int: if type(value) is not int or value < 0 or value > (2**63 - 1): raise ProtocolError( "invalid_event_seq", "event_seq must be an integer from 0 through 2^63-1", ) return value def _optional_unix_time(value: Any) -> float | None: if value is None: return None if isinstance(value, bool) or not isinstance(value, (int, float)): raise ProtocolError( "invalid_sent_unix_s", "sent_unix_s must be a finite number" ) result = float(value) if not math.isfinite(result) or result < 0.0: raise ProtocolError( "invalid_sent_unix_s", "sent_unix_s must be a finite number" ) return result def parse_control_message(text: str) -> dict[str, Any]: """Parse one strict but forward-compatible recorder control object.""" if not isinstance(text, str) or not text or len(text.encode("utf-8")) > 16384: raise ProtocolError( "invalid_json", "control payload must be 1 through 16384 UTF-8 bytes" ) try: raw = json.loads(text) except (json.JSONDecodeError, UnicodeError) as error: raise ProtocolError("invalid_json", f"invalid JSON: {error}") from error if not isinstance(raw, Mapping): raise ProtocolError("invalid_json_object", "control JSON must be an object") if raw.get("version") != PROTOCOL_VERSION: raise ProtocolError( "unsupported_version", f"version must equal {PROTOCOL_VERSION}" ) command = raw.get("command") if command not in ("start", "stop", "heartbeat"): raise ProtocolError( "invalid_command", "command must be start, stop, or heartbeat" ) event_seq = _event_sequence(raw.get("event_seq")) request_id = _nonempty_text(raw.get("request_id"), "request_id") session_id = _nonempty_text( raw.get("teleop_session_id"), "teleop_session_id" ) capture_id = _capture_id(raw.get("capture_id")) sent_unix_s = _optional_unix_time(raw.get("sent_unix_s")) reason_value = raw.get("reason", "bridge_heartbeat" if command == "heartbeat" else "button") reason = _nonempty_text(reason_value, "reason", maximum=512) return { "version": PROTOCOL_VERSION, "command": command, "event_seq": event_seq, "request_id": request_id, "teleop_session_id": session_id, "capture_id": capture_id, "reason": reason, "sent_unix_s": sent_unix_s, } def topics_without_publishers( required_topics: Sequence[str], publisher_lookup: Callable[[str], Sequence[Any]], ) -> tuple[str, ...]: """Return required topics that currently have no live publisher endpoint. ROS graph discovery can retain a topic name solely because this process is subscribed to it. Looking up publisher endpoints is therefore the meaningful start preflight; topic-name presence alone is insufficient. Exceptions deliberately propagate so the caller can report graph-query failure separately from an ordinary missing publisher. """ missing: list[str] = [] for topic in required_topics: if not publisher_lookup(topic): missing.append(topic) return tuple(missing) class RecorderControlProtocol: """Idempotent bridge-to-recorder protocol and heartbeat watchdog.""" def __init__( self, recorder: RecorderBackend, *, heartbeat_timeout_seconds: float = 3.0, monotonic: Callable[[], float] = time.monotonic, unix_time: Callable[[], float] = time.time, start_preflight: Callable[[], tuple[bool, str]] | None = None, request_cache_size: int = 512, session_cache_size: int = 64, ) -> None: if ( not math.isfinite(heartbeat_timeout_seconds) or heartbeat_timeout_seconds <= 0.0 ): raise ValueError("heartbeat timeout must be positive and finite") if request_cache_size <= 0 or session_cache_size <= 0: raise ValueError("protocol cache sizes must be positive") self.recorder = recorder self.heartbeat_timeout_seconds = float(heartbeat_timeout_seconds) self._monotonic = monotonic self._unix_time = unix_time self._start_preflight = start_preflight self._request_cache_size = request_cache_size self._session_cache_size = session_cache_size self._request_cache: OrderedDict[str, tuple[str, bool, str, int]] = ( OrderedDict() ) self._last_sequence_by_session: OrderedDict[str, int] = OrderedDict() self._capture_id: str | None = None self._session_id: str | None = None self._last_heartbeat_monotonic: float | None = None self._watchdog_stop_requested = False self._last_protocol_error = "" self._ack_event_seq: int | None = None self._ack_request_id: str | None = None self._ack_accepted: bool | None = None self._ack_code = "initialized" self._accepted_count = 0 self._rejected_count = 0 self._duplicate_count = 0 self._watchdog_stop_count = 0 def handle_json(self, text: str, *, now: float | None = None) -> dict[str, Any]: """Handle one payload without blocking on recorder I/O.""" timestamp = self._monotonic() if now is None else float(now) try: event = parse_control_message(text) except ProtocolError as error: self._ack_event_seq = None self._ack_request_id = None self._ack_accepted = False self._ack_code = error.code self._last_protocol_error = str(error) self._rejected_count += 1 return self.status(now=timestamp) fingerprint = json.dumps(event, sort_keys=True, separators=(",", ":")) request_id = event["request_id"] cached = self._request_cache.get(request_id) if cached is not None: cached_fingerprint, accepted, code, event_seq = cached if cached_fingerprint != fingerprint: return self._reject( event, "request_id_reused", "request_id was already used for a different payload", fingerprint=fingerprint, cache=False, now=timestamp, ) self._request_cache.move_to_end(request_id) self._duplicate_count += 1 self._set_ack(event_seq, request_id, accepted, code) return self.status(now=timestamp) session_id = event["teleop_session_id"] last_sequence = self._last_sequence_by_session.get(session_id) if last_sequence is not None and event["event_seq"] <= last_sequence: return self._reject( event, "stale_event_seq", f"event_seq must be greater than the previous value {last_sequence}", fingerprint=fingerprint, now=timestamp, ) self._remember_sequence(session_id, event["event_seq"]) self._last_protocol_error = "" try: accepted, code = self._execute(event, timestamp) message = ( "" if accepted else (self._last_protocol_error or code) ) except Exception as error: accepted = False code = "recorder_exception" message = f"{type(error).__name__}: {error}" if accepted: self._accepted_count += 1 self._last_protocol_error = "" else: self._rejected_count += 1 self._last_protocol_error = message self._remember_request( request_id, fingerprint, accepted, code, event["event_seq"], ) self._set_ack( event["event_seq"], request_id, accepted, code, ) return self.status(now=timestamp) def poll(self, *, now: float | None = None) -> dict[str, Any]: """Advance the watchdog and return the current status snapshot.""" timestamp = self._monotonic() if now is None else float(now) recorder_status = self._safe_recorder_status() manager_state = str(recorder_status.get("state", "unknown")) if ( manager_state in ("start_pending", "starting", "recording") and self._last_heartbeat_monotonic is not None and not self._watchdog_stop_requested and timestamp - self._last_heartbeat_monotonic > self.heartbeat_timeout_seconds ): try: accepted = self.recorder.request_stop("bridge_heartbeat_timeout") except Exception as error: self._last_protocol_error = ( f"watchdog stop failed: {type(error).__name__}: {error}" ) else: if accepted: self._watchdog_stop_requested = True self._watchdog_stop_count += 1 return self._status_from_snapshot(recorder_status, timestamp) def status(self, *, now: float | None = None) -> dict[str, Any]: timestamp = self._monotonic() if now is None else float(now) return self._status_from_snapshot(self._safe_recorder_status(), timestamp) def shutdown(self, *, timeout: float | None = None) -> bool: """Request clean recorder finalization when the supervisor exits.""" return self.recorder.shutdown(wait=True, timeout=timeout) def _execute(self, event: Mapping[str, Any], now: float) -> tuple[bool, str]: command = event["command"] if command == "start": return self._start(event, now) if command == "stop": return self._stop(event) return self._heartbeat(event, now) def _start(self, event: Mapping[str, Any], now: float) -> tuple[bool, str]: recorder_status = self._safe_recorder_status() manager_state = str(recorder_status.get("state", "unknown")) if manager_state in ACTIVE_MANAGER_STATES: if self._matches_context(event): return True, "already_active" return False, "recorder_busy" if manager_state not in ("idle",): return False, "recorder_unavailable" if self._start_preflight is not None: allowed, detail = self._start_preflight() if not allowed: self._last_protocol_error = detail return False, "preflight_failed" actual_capture_id = self.recorder.request_start( event["teleop_session_id"], event["reason"], episode_id=event["capture_id"], ) if actual_capture_id is None: return False, "recorder_busy" if actual_capture_id != event["capture_id"]: # A mismatched directory would make bridge samples and manifests # impossible to correlate, so immediately fail safe by stopping. self.recorder.request_stop("capture_id_mismatch") return False, "capture_id_mismatch" self._capture_id = actual_capture_id self._session_id = event["teleop_session_id"] self._last_heartbeat_monotonic = now self._watchdog_stop_requested = False return True, "start_accepted" def _stop(self, event: Mapping[str, Any]) -> tuple[bool, str]: recorder_status = self._safe_recorder_status() manager_state = str(recorder_status.get("state", "unknown")) if manager_state not in ACTIVE_MANAGER_STATES: if self._matches_context(event) or self._capture_id is None: return True, "already_stopped" return False, "capture_not_active" if not self._matches_context(event): return False, "capture_mismatch" accepted = self.recorder.request_stop(event["reason"]) if not accepted: return False, "stop_rejected" self._watchdog_stop_requested = True return True, "stop_accepted" def _heartbeat( self, event: Mapping[str, Any], now: float ) -> tuple[bool, str]: if not self._matches_context(event): return False, "capture_mismatch" recorder_status = self._safe_recorder_status() manager_state = str(recorder_status.get("state", "unknown")) if manager_state not in ACTIVE_MANAGER_STATES: return True, "already_stopped" self._last_heartbeat_monotonic = now return True, "heartbeat_accepted" def _matches_context(self, event: Mapping[str, Any]) -> bool: return ( self._capture_id == event["capture_id"] and self._session_id == event["teleop_session_id"] ) def _reject( self, event: Mapping[str, Any], code: str, message: str, *, fingerprint: str, cache: bool = True, now: float, ) -> dict[str, Any]: self._rejected_count += 1 self._last_protocol_error = message self._set_ack( event["event_seq"], event["request_id"], False, code ) if cache: self._remember_request( event["request_id"], fingerprint, False, code, event["event_seq"], ) return self.status(now=now) def _remember_sequence(self, session_id: str, sequence: int) -> None: self._last_sequence_by_session[session_id] = sequence self._last_sequence_by_session.move_to_end(session_id) while len(self._last_sequence_by_session) > self._session_cache_size: self._last_sequence_by_session.popitem(last=False) def _remember_request( self, request_id: str, fingerprint: str, accepted: bool, code: str, event_seq: int, ) -> None: self._request_cache[request_id] = ( fingerprint, accepted, code, event_seq, ) self._request_cache.move_to_end(request_id) while len(self._request_cache) > self._request_cache_size: self._request_cache.popitem(last=False) def _set_ack( self, event_seq: int, request_id: str, accepted: bool, code: str, ) -> None: self._ack_event_seq = event_seq self._ack_request_id = request_id self._ack_accepted = accepted self._ack_code = code def _safe_recorder_status(self) -> dict[str, Any]: try: status = self.recorder.status() except Exception as error: self._last_protocol_error = ( f"recorder status failed: {type(error).__name__}: {error}" ) return { "state": "unknown", "recording": False, "last_error": self._last_protocol_error, } return dict(status) def _status_from_snapshot( self, recorder_status: Mapping[str, Any], now: float ) -> dict[str, Any]: manager_state = str(recorder_status.get("state", "unknown")) if manager_state in ("start_pending", "starting"): state = "starting" elif manager_state == "recording": state = "recording" elif manager_state in ("stop_pending", "finalizing"): state = "stopping" elif manager_state == "idle": last_result = recorder_status.get("last_result") state = last_result if last_result in ("ready", "failed") else "idle" else: state = "failed" heartbeat_age: float | None = None if self._last_heartbeat_monotonic is not None: heartbeat_age = max(0.0, now - self._last_heartbeat_monotonic) manager_error = recorder_status.get("last_error") last_error = ( str(manager_error) if isinstance(manager_error, str) and manager_error else self._last_protocol_error ) return { "version": PROTOCOL_VERSION, "state": state, "capture_id": self._capture_id, "teleop_session_id": self._session_id, "ack_event_seq": self._ack_event_seq, "ack_request_id": self._ack_request_id, "ack_accepted": self._ack_accepted, "ack_code": self._ack_code, "last_error": last_error, "stop_reason": recorder_status.get("stop_reason"), "heartbeat_age_s": heartbeat_age, "heartbeat_timeout_s": self.heartbeat_timeout_seconds, "updated_unix_s": float(self._unix_time()), "statistics": { "accepted": self._accepted_count, "rejected": self._rejected_count, "duplicates": self._duplicate_count, "watchdog_stops": self._watchdog_stop_count, }, "recorder": dict(recorder_status), } __all__ = [ "PROTOCOL_VERSION", "ProtocolError", "RecorderControlProtocol", "parse_control_message", "topics_without_publishers", ]