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