dmitchelljackson's picture
Upload Qwen3.5 AndroidWorld RL harness milestone 2026-06-09
b7df156 verified
Raw History Blame Contribute Delete
9.84 kB
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))