File size: 6,009 Bytes
b7df156
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
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")