from __future__ import annotations import base64 import json import subprocess import time import urllib.request from io import BytesIO from typing import Any from PIL import Image from data.datasets.android_control_a11y import A11yNode, A11yWindow from data.datasets.android_control_a11y import compress_windowed_tree, expand_range_targets, get_action_targets from rl_harness.types import HarnessFrame class ApkStateCollector: """Model-facing state source from the Cerebellum Android collector APK.""" def __init__( self, adb_path: str = "/opt/android/platform-tools/adb", serial: str | None = None, host: str = "127.0.0.1", port: int = 8765, device_port: int = 8765, max_tree_chars: int = 5000, retries: int = 6, retry_sleep: float = 0.5, ): self.adb_path = adb_path self.serial = serial self.host = host self.port = port self.device_port = device_port self.max_tree_chars = max_tree_chars self.retries = retries self.retry_sleep = retry_sleep def _adb_cmd(self, *args: str) -> list[str]: cmd = [self.adb_path] if self.serial: cmd.extend(["-s", self.serial]) else: cmd.append("-e") cmd.extend(args) return cmd def ensure_forward(self) -> None: subprocess.run( self._adb_cmd("forward", f"tcp:{self.port}", f"tcp:{self.device_port}"), stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=False, timeout=10.0, ) def restart_service(self) -> None: service = "com.cerebellum.collector/.CollectorAccessibilityService" current = subprocess.run( self._adb_cmd("shell", "settings", "get", "secure", "enabled_accessibility_services"), stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=False, timeout=10.0, ).stdout.strip() services = [s for s in current.split(":") if s and s != "null"] if service not in services: services.append(service) subprocess.run( self._adb_cmd( "shell", "settings", "put", "secure", "enabled_accessibility_services", ":".join(services), ), stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=False, timeout=10.0, ) subprocess.run( self._adb_cmd("shell", "settings", "put", "secure", "accessibility_enabled", "1"), stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=False, timeout=10.0, ) self.ensure_forward() def _fetch_frame(self) -> dict[str, Any]: self.ensure_forward() with urllib.request.urlopen(f"http://{self.host}:{self.port}/frame", timeout=5.0) as resp: return json.loads(resp.read().decode("utf-8")) def _decode_image(self, row: dict[str, Any]) -> Image.Image: b64 = row.get("screenshot_png_b64") if not b64: raise RuntimeError("collector APK returned no screenshot") return Image.open(BytesIO(base64.b64decode(b64))).convert("RGB") def _parse_node(self, row: dict[str, Any], nodes: list[A11yNode]) -> A11yNode: bounds = row.get("bounds") node = A11yNode() node.node_id = len(nodes) node.bounds = tuple(int(x) for x in bounds) if bounds and len(bounds) == 4 else None node.text = row.get("text") or "" node.content_desc = row.get("content_desc") or "" node.resource_id = row.get("view_id") or "" node.class_name = row.get("class_name") or "" node.visible = bool(row.get("visible", True)) node.is_enabled = bool(row.get("enabled", True)) node.is_focused = bool(row.get("focused", False)) node.is_editable = bool(row.get("editable", False)) node.is_scrollable = bool(row.get("scrollable", False)) node.is_clickable = bool(row.get("clickable", False)) node.is_long_clickable = bool(row.get("long_clickable", False)) node.is_checkable = bool(row.get("checkable", False)) node.is_checked = bool(row.get("checked", False)) node.is_selected = bool(row.get("selected", False)) node.is_password = bool(row.get("password", False)) range_info = row.get("range") or {} if range_info: node.range_current = range_info.get("current") node.range_min = range_info.get("min") node.range_max = range_info.get("max") node.range_type = range_info.get("type") node.action_ids = [int(x) for x in row.get("actions") or [] if isinstance(x, int)] # SFT parser sets these from the accessibility action list. The APK # service exposes the equivalent high-level flags; use them as the # action-space source so get_action_targets() sees the same fields. node.clickable = node.is_clickable node.long_clickable = node.is_long_clickable node.scrollable = node.is_scrollable nodes.append(node) for child in row.get("children") or []: child_node = self._parse_node(child, nodes) node.child_ids.append(child_node.node_id) node.children.append(child_node) return node def _package_candidates(self, row: dict[str, Any]) -> set[str]: packages = set() def walk(node: dict[str, Any]) -> None: package_name = node.get("package") or node.get("package_name") if package_name: packages.add(str(package_name)) resource_id = node.get("view_id") or "" if ":id/" in resource_id: packages.add(resource_id.split(":id/", 1)[0]) for child in node.get("children") or []: walk(child) for window in row.get("windows") or []: root = window.get("root") if root: walk(root) return packages def _windows(self, row: dict[str, Any]) -> list[A11yWindow]: windows: list[A11yWindow] = [] for window in row.get("windows") or []: nodes: list[A11yNode] = [] root = window.get("root") if root: self._parse_node(root, nodes) windows.append( A11yWindow( idx=int(window.get("id", len(windows))), title=window.get("title") or "", type_=window.get("type"), layer=window.get("layer"), active=window.get("active"), focused=window.get("focused"), nodes=nodes, ) ) return windows def capture_frame(self) -> HarnessFrame: last_error: Exception | None = None for attempt in range(self.retries): try: row = self._fetch_frame() image = self._decode_image(row) windows = self._windows(row) expand_range_targets(windows, zones=5) tappable = get_action_targets(windows) tree = compress_windowed_tree(windows, tappable) if self.max_tree_chars and len(tree) > self.max_tree_chars: tree = tree[: self.max_tree_chars].rstrip() + "\n... tree truncated ..." return HarnessFrame( image=image, tappable=tappable, compressed_tree=tree, packages=sorted(self._package_candidates(row)), ) except Exception as exc: last_error = exc if attempt >= 1: self.restart_service() time.sleep(self.retry_sleep) raise RuntimeError(f"APK state capture failed: {last_error}") from last_error @staticmethod def _frame_signature(frame: HarnessFrame) -> tuple: targets = [] for node in frame.tappable: targets.append(( getattr(node, "bounds", None), getattr(node, "text", "") or "", getattr(node, "content_desc", "") or "", getattr(node, "resource_id", "") or "", getattr(node, "class_name", "") or "", )) return (frame.image.size, frame.compressed_tree, tuple(targets)) def capture_stable_frame( self, max_wait_s: float = 1.0, interval_s: float = 0.2, min_captures: int = 2, min_wait_s: float = 0.5, ) -> HarnessFrame: """Capture after the APK tree has stopped changing. Accessibility nodes often arrive a few hundred milliseconds after the screenshot, especially immediately after opening an app. Returning that partial tree gives the model a different label space from SFT. This bounded poll waits for the model-facing frame signature to repeat. """ start = time.time() deadline = start + max(0.0, max_wait_s) previous_sig = None last_frame = None captures = 0 while True: frame = self.capture_frame() captures += 1 sig = self._frame_signature(frame) if ( captures >= min_captures and previous_sig == sig and time.time() - start >= max(0.0, min_wait_s) ): return frame previous_sig = sig last_frame = frame if time.time() >= deadline: return last_frame time.sleep(max(0.01, interval_s))