from __future__ import annotations from collections import deque from dataclasses import dataclass, field from typing import Any from rl_harness.types import ParsedAction NORMAL_POST_ACTION_TIMEOUT_MS = 1000 TERMINAL_CAPTURE_TIMEOUT_MS = 3000 @dataclass(frozen=True) class LabelFingerprint: bounds: tuple[int, int, int, int] | None text: str = "" content_desc: str = "" resource_id: str = "" class_name: str = "" enabled: bool = True clickable: bool = False long_clickable: bool = False focused: bool = False window_id: str = "" window_type: str = "" @dataclass class StreamFrame: seq: int timestamp_ms: int image: Any = None tree: str = "" label_map: dict[str, LabelFingerprint] = field(default_factory=dict) focused: LabelFingerprint | None = None app_id: str = "" window_id: str = "" event_summary: str = "" @dataclass class PendingDecision: source_seq: int source_frame: StreamFrame started_at_ms: int @dataclass class ValidationResult: ok: bool reason: str class FrameBuffer: """Python-side rolling observation queue. The APK/bridge should stream raw frames. This object owns the model-facing causal bookkeeping: current frame selection, stale-action validation, and frame collection after an action. """ def __init__(self, maxlen: int = 64): self.frames: deque[StreamFrame] = deque(maxlen=maxlen) def push(self, frame: StreamFrame) -> None: if self.frames and frame.seq <= self.frames[-1].seq: return self.frames.append(frame) @property def latest(self) -> StreamFrame | None: return self.frames[-1] if self.frames else None def after(self, seq: int) -> list[StreamFrame]: return [f for f in self.frames if f.seq > seq] def transition_frames(self, source_seq: int) -> list[StreamFrame]: return self.after(source_seq) def fingerprint_from_node(node: Any, window_id: str = "", window_type: str = "") -> LabelFingerprint: bounds = getattr(node, "bounds", None) if bounds is not None: bounds = tuple(int(x) for x in bounds) return LabelFingerprint( bounds=bounds, text=getattr(node, "text", "") or "", content_desc=getattr(node, "content_desc", "") or "", resource_id=getattr(node, "resource_id", "") or "", class_name=getattr(node, "class_name", "") or "", enabled=bool(getattr(node, "is_enabled", True)), clickable=bool(getattr(node, "clickable", False) or getattr(node, "is_clickable", False)), long_clickable=bool(getattr(node, "long_clickable", False) or getattr(node, "is_long_clickable", False)), focused=bool(getattr(node, "is_focused", False)), window_id=str(window_id or ""), window_type=str(window_type or ""), ) def label_fingerprints_from_sample(sample: dict, label_map: dict[int, str]) -> dict[str, LabelFingerprint]: out = {} tappable = sample.get("tappable", []) or [] for idx, label in label_map.items(): if idx < 0 or idx >= len(tappable): continue out[label] = fingerprint_from_node(tappable[idx]) return out def focused_fingerprint_from_sample(sample: dict) -> LabelFingerprint | None: nodes = sample.get("all_nodes", []) or [] focused = [ n for n in nodes if getattr(n, "visible", False) and getattr(n, "is_focused", False) and getattr(n, "is_editable", False) and getattr(n, "bounds", None) ] if not focused: return None node = min(focused, key=lambda n: (n.bounds[2] - n.bounds[0]) * (n.bounds[3] - n.bounds[1])) return fingerprint_from_node(node) def stream_frame_from_policy_sample( sample: dict, label_map: dict[int, str], seq: int, timestamp_ms: int, app_id: str = "", window_id: str = "", ) -> StreamFrame: return StreamFrame( seq=seq, timestamp_ms=timestamp_ms, image=sample.get("current_img"), tree=sample.get("compressed_tree", ""), label_map=label_fingerprints_from_sample(sample, label_map), focused=focused_fingerprint_from_sample(sample), app_id=app_id, window_id=window_id, ) def validate_action_against_latest( action: ParsedAction, source_frame: StreamFrame, latest_frame: StreamFrame, ) -> ValidationResult: """Decide whether a completed prediction is still executable. The action is tied to `source_frame`; only exact selected-label/focus drift invalidates it. Spinner/progress-only frames can advance `seq` without causing rejection. """ if action.kind in {"tap", "long_press"}: if not action.label: return ValidationResult(False, "missing_label") source_fp = source_frame.label_map.get(action.label) latest_fp = latest_frame.label_map.get(action.label) if source_fp is None: return ValidationResult(False, "source_label_missing") if latest_fp is None: return ValidationResult(False, "latest_label_missing") if latest_fp != source_fp: return ValidationResult(False, "label_fingerprint_changed") return ValidationResult(True, "same_label_fingerprint") if action.kind == "type": if source_frame.focused is None: return ValidationResult(False, "source_focus_missing") if latest_frame.focused != source_frame.focused: return ValidationResult(False, "focused_fingerprint_changed") return ValidationResult(True, "same_focused_fingerprint") if source_frame.app_id and latest_frame.app_id and source_frame.app_id != latest_frame.app_id: return ValidationResult(False, "app_changed") if source_frame.window_id and latest_frame.window_id and source_frame.window_id != latest_frame.window_id: return ValidationResult(False, "window_changed") return ValidationResult(True, "non_label_action_ok")