File size: 7,221 Bytes
ed79a7f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
"""
Graph Machine DOM Referral Engine.

Synthesized from:
"Graph Machine: Towards Better Pretraining via Edges" (Iter Labs, Sep 2, 2026, arXiv:2609.02881).

Represents the DOM as an O(n) state graph with pointer-like directed edges and weights.
Uses dynamic 2-hop pointer chasing (referral routing) to address interactive elements
in O(1) retrieval hops rather than linearizing and scanning massive flat DOM dumps.
"""

from __future__ import annotations

import math
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Set, Tuple
from miniswardbower.core.schemas import InteractiveElement, PrunedAXTree


@dataclass
class ReferralEdge:
    """Directed pointer edge from source node to target node with referral weight."""
    target_id: str
    edge_type: str  # 'parent', 'child', 'sibling', 'semantic_label', 'spatial_neighbor'
    weight: float = 1.0


@dataclass
class DOMGraphNode:
    """Node in the Graph Machine DOM representation."""
    node_id: str
    tag: str
    role: str
    text: str
    bbox: Optional[Tuple[float, float, float, float]] = None
    edges: List[ReferralEdge] = field(default_factory=list)


class DOMReferralGraph:
    """
    Graph Machine DOM representation with dynamic pointer referral routing.
    Enables O(1) element lookup via 2-hop referral traversal.
    """

    def __init__(self, max_edges_per_node: int = 8):
        self.max_edges = max_edges_per_node
        self.nodes: Dict[str, DOMGraphNode] = {}
        self.root_id: Optional[str] = None

    @classmethod
    def from_pruned_tree(cls, tree: PrunedAXTree) -> DOMReferralGraph:
        """Constructs a DOM referral graph from a pruned accessibility tree."""
        graph = cls()

        # 1. Create nodes
        for el in tree.elements:
            node = DOMGraphNode(
                node_id=el.id,
                tag=el.tag,
                role=el.element_type or el.tag,
                text=el.text or el.placeholder or el.aria_label or "",
                bbox=el.bbox,
            )
            graph.nodes[el.id] = node
            if graph.root_id is None:
                graph.root_id = el.id

        # 2. Add structural & spatial edges
        el_list = list(tree.elements)
        n = len(el_list)
        for i in range(n):
            src = el_list[i]
            src_node = graph.nodes[src.id]

            # Sequential sibling edges
            if i > 0:
                src_node.edges.append(ReferralEdge(target_id=el_list[i - 1].id, edge_type="prev_sibling", weight=0.6))
            if i < n - 1:
                src_node.edges.append(ReferralEdge(target_id=el_list[i + 1].id, edge_type="next_sibling", weight=0.6))

            # Spatial 2D proximity edges (pointer to nearest visual neighbor)
            if src.bbox:
                sx, sy, sw, sh = src.bbox
                scx, scy = sx + sw / 2.0, sy + sh / 2.0

                min_dist = float("inf")
                nearest_id = None
                for j in range(n):
                    if i == j:
                        continue
                    tgt = el_list[j]
                    if tgt.bbox:
                        tx, ty, tw, th = tgt.bbox
                        tcx, tcy = tx + tw / 2.0, ty + th / 2.0
                        dist = math.sqrt((scx - tcx) ** 2 + (scy - tcy) ** 2)
                        if dist < min_dist:
                            min_dist = dist
                            nearest_id = tgt.id

                if nearest_id:
                    spatial_weight = 1.0 / (1.0 + min_dist / 100.0)
                    src_node.edges.append(ReferralEdge(target_id=nearest_id, edge_type="spatial_neighbor", weight=spatial_weight))

            # Semantic edges: associate input fields with nearby text labels
            if src.tag in ("input", "textarea", "select"):
                for j in range(max(0, i - 3), i):
                    prev_el = el_list[j]
                    if prev_el.text:
                        src_node.edges.append(ReferralEdge(target_id=prev_el.id, edge_type="semantic_label", weight=0.9))

        return graph

    def referral_route(
        self,
        query: str,
        start_node_id: Optional[str] = None,
        max_hops: int = 4,
    ) -> Tuple[Optional[str], float, List[str]]:
        """
        Executes pointer referral chasing to find the most relevant element node.
        Uses GM dynamic pointer addressing: seeds entry from token pointers then
        performs 2-hop referral pointer chasing along local neighborhood edges.
        Returns: (target_node_id, match_score, referral_path)
        """
        query_terms = set(query.lower().split())
        if not self.nodes:
            return None, 0.0, []

        # 1. Pointer referral seed: find entry node containing key token pointers
        if start_node_id and start_node_id in self.nodes:
            start_id = start_node_id
        else:
            # Seed entry point via pointer referral
            best_seed_id = self.root_id
            best_seed_score = -1.0
            for nid, node in self.nodes.items():
                score = self._compute_relevance(node, query_terms)
                if score > best_seed_score:
                    best_seed_score = score
                    best_seed_id = nid
            start_id = best_seed_id

        current_id = start_id
        path = [current_id]
        best_match_id = current_id
        best_score = self._compute_relevance(self.nodes[current_id], query_terms)

        visited: Set[str] = {current_id}

        for hop in range(max_hops):
            curr_node = self.nodes.get(current_id)
            if not curr_node:
                break
                break

            # 2-hop referral: evaluate outgoing edges and their referral targets
            next_id = None
            highest_referral_val = -1.0

            for edge in curr_node.edges:
                tgt_id = edge.target_id
                if tgt_id in visited or tgt_id not in self.nodes:
                    continue

                tgt_node = self.nodes[tgt_id]
                rel = self._compute_relevance(tgt_node, query_terms)
                val = rel * edge.weight

                if rel > best_score:
                    best_score = rel
                    best_match_id = tgt_id

                if val > highest_referral_val:
                    highest_referral_val = val
                    next_id = tgt_id

            if next_id is None or highest_referral_val <= 0.0:
                break

            current_id = next_id
            visited.add(current_id)
            path.append(current_id)

            if best_score >= 0.85:
                # Early stop on high-confidence match
                break

        return best_match_id, best_score, path

    def _compute_relevance(self, node: DOMGraphNode, query_terms: Set[str]) -> float:
        """Computes overlap relevance between node semantics and query terms."""
        content = f"{node.tag} {node.role} {node.text}".lower()
        node_words = set(content.split())
        if not query_terms or not node_words:
            return 0.0

        overlap = query_terms.intersection(node_words)
        return len(overlap) / len(query_terms)