File size: 14,856 Bytes
adf912b
 
 
 
 
 
454b3e6
adf912b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
454b3e6
 
 
 
 
adf912b
 
454b3e6
 
 
 
 
 
 
 
 
 
 
 
adf912b
 
454b3e6
adf912b
 
 
 
 
 
 
 
 
454b3e6
 
 
 
 
 
 
 
adf912b
454b3e6
adf912b
 
 
 
 
454b3e6
 
 
adf912b
 
 
 
 
 
 
 
 
 
454b3e6
 
adf912b
 
 
ac29aef
 
 
 
 
 
 
 
 
 
 
 
 
adf912b
 
 
 
454b3e6
 
 
 
 
 
 
 
adf912b
 
 
 
 
 
 
 
 
 
 
454b3e6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
adf912b
 
 
 
 
 
 
 
454b3e6
 
 
 
 
 
 
 
 
adf912b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
454b3e6
adf912b
 
454b3e6
 
adf912b
454b3e6
adf912b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
"""TypeSafe-compatible /v1/systemone endpoint backed by a local laya checkpoint (+ TileLang fast path).

    python apps/systemone_server.py [port=8791] [variant=typed|multilingual|english]

Request body: {"model": ..., "state": {...}, "questions": {...}}  ->  {"answers": ..., "model": ..., "usage": ...}
"""
import json, os, re, sys, time, traceback
from http.server import ThreadingHTTPServer, BaseHTTPRequestHandler
from common import get_agent
from fast_batch import predict_fast

PORT = int(sys.argv[1]) if len(sys.argv) > 1 else 8791
VARIANT = sys.argv[2] if len(sys.argv) > 2 else "typed"
MAXOPT = int(sys.argv[3]) if len(sys.argv) > 3 else 12   # laya's option budget: keep choice questions at most this wide
agent = get_agent(VARIANT)
FMT = os.environ.get("LAYA_FMT", agent.cfg.get("laya_fmt", "v1"))   # fine-tuned checkpoints record their format in rl_agent_config.json
if agent.cfg.get("head_max_len_train"):
    agent.cfg["head_max_len"] = agent.cfg["head_max_len_train"]
print(f"[systemone] format={FMT} head_max_len={agent.cfg.get('head_max_len')}", flush=True)
LOG = []

# ---- System 1 / System 2 gating: below ESCALATE_TAU confidence, ask the LLM teacher (same element table) and return its
# decision in laya's answer format.  Every escalation is also appended to ESCALATE_LOG as a DAgger case.
TAU = float(os.environ.get("ESCALATE_TAU", "0"))           # 0 = off
# System 2 model: S2_BASE_URL / S2_API_KEY / S2_MODEL / S2_EXTRA_JSON (e.g. DeepSeek); defaults to the local text model
ESC_URL = os.environ.get("S2_BASE_URL", os.environ.get("TEXT_MODEL_BASE_URL", "http://127.0.0.1:30000/v1")).rstrip("/") + "/chat/completions"
ESC_MODEL = os.environ.get("S2_MODEL", os.environ.get("TEXT_MODEL", "Qwen/Qwen3-8B-AWQ"))
ESC_KEY = os.environ.get("S2_API_KEY", "")
ESC_EXTRA = json.loads(os.environ.get("S2_EXTRA_JSON", '{"chat_template_kwargs": {"enable_thinking": false}}'))
ESC_LOG = os.environ.get("ESCALATE_LOG", "")
STATS = {"calls": 0, "escalated": 0}
ESC_SYS = """You are the careful System-2 decision maker of a browser agent. You get the user's goal, the recent actions
(with whether each changed the page), the current page (url, title, visible text), the available OPERATIONS (key ->
meaning) and, per operation, the numbered target controls. Choose the single best NEXT step toward the WHOLE goal.
Rules:
- Think first in "thought" (1-3 sentences): what is already done, what is missing, which control does it.
- After typing a search/query, SUBMIT it: PRESS_ENTER (if offered) or click the search button / the matching suggestion.
- WAIT only when results are visibly loading; never WAIT twice in a row. If the last actions changed nothing, do something
  different (another control, scroll, open a menu/filter).
- Apply every requested filter/sort/value; open the requested item. DONE only when every requirement is visibly met.
- Close cookie/consent/newsletter popups only if they block the page. Never log in, pay, order or send messages.
- BLOCKED only if no offered operation can make progress.
Answer JSON: {"thought": "...", "operation": "<one OPERATIONS key>", "target": "<target key for that operation, or null>"}"""


def escalate(state, questions, answers, reason=None):
    import httpx
    ops = questions["operation"]["criteria"]
    controls = []
    for qid, q in questions.items():
        if qid.endswith("_target"):
            op = qid[:-7].upper()
            for key, desc in q["criteria"].items():
                controls.append({"op": op, "target": key, "control": desc})
    user = {"goal": (questions["operation"]["instructions"] or {}).get("goal") if isinstance(questions["operation"]["instructions"], dict) else "",
            "recent_actions": state.get("recent_actions", [])[-8:], "page": state.get("page", {}),
            "OPERATIONS": {k: (v if isinstance(v, str) else str(v)) for k, v in ops.items()} if isinstance(ops, dict) else list(ops),
            "targets": {op: {c["target"]: c["control"] for c in controls if c["op"] == op} for op in sorted({c["op"] for c in controls})},
            "fast_policy_guess": {k: v.get("choice") for k, v in answers.items()}}
    if reason:
        user["why_you_are_asked"] = {"done_rejected": "the fast policy said DONE but a checker found the task NOT complete yet: find the missing part",
                                     "stuck": "the fast policy's last actions changed nothing or repeat: choose a different, useful action"}.get(reason, reason)
    body = {"model": ESC_MODEL, "max_tokens": 400, "temperature": 0.0, "response_format": {"type": "json_object"}, **ESC_EXTRA,
            "messages": [{"role": "system", "content": ESC_SYS}, {"role": "user", "content": json.dumps(user, ensure_ascii=False)}]}
    r = httpx.post(ESC_URL, json=body, timeout=120, headers={"Authorization": f"Bearer {ESC_KEY}"} if ESC_KEY else None).json()
    v = json.loads(r["choices"][0]["message"]["content"])
    op = str(v.get("operation", "")).upper(); tgt = v.get("target")
    if op not in ops:
        return answers, False
    def one_hot(keys, k, p=0.97):
        if len(keys) == 1:          # a single option must carry all the mass (the client checks that they sum to 1)
            return {k: 1.0}
        rest = (1 - p) / (len(keys) - 1)
        return {kk: (p if kk == k else rest) for kk in keys}
    answers["operation"] = {**answers["operation"], "choice": op, "probabilities": one_hot(list(ops), op), "confidence": 0.9, "system2": True}
    tq = op.lower() + "_target"
    if tq in questions:
        keys = list(questions[tq]["criteria"]); tgt = str(tgt)
        if tgt not in keys:
            return answers, False
        answers[tq] = {**answers[tq], "choice": tgt, "probabilities": one_hot(keys, tgt), "confidence": 0.9, "system2": True}
    if ESC_LOG:
        with open(ESC_LOG, "a") as f:
            f.write(json.dumps({"state": state, "questions": questions, "reason": reason, "system1": {k: v.get("choice") for k, v in answers.items()},
                                "operation": op, "target": tgt if tq in questions else None}, ensure_ascii=False) + "\n")
    return answers, True


def _cut(el, n):
    """Truncate an element label to n chars; for a <select> option ("[i:k] Field label → Option") keep the option name."""
    el = str(el)
    if len(el) <= n:
        return el
    if " → " in el:
        head, opt = el.rsplit(" → ", 1)
        opt = " ".join(opt.split())[:40]
        keep = max(12, n - len(opt) - 3)
        return " ".join(head.split())[:keep] + " → " + opt
    return el[:n]


def compact(v):
    """Shrink jev-ultrafast element criteria ({'element': '[3] Search', 'role': 'button', ...}) into one short string
    so more options fit laya's head token budget."""
    if isinstance(v, dict) and "element" in v:
        el = str(v["element"])
        if FMT in ("v4", "v5", "v6"):
            # v4: the option key is already rendered by laya ("<key>: ..."), so drop the duplicate "[key] "; a <select>
            # option is only "Field → Option" (role and current value repeated on every option ate the head budget)
            el = re.sub(r"^\[[^\]]*\]\s*", "", el)
            if " → " in el:
                return _cut(el, 50)
        s = _cut(el, 50 if FMT in ("v3", "v4", "v5", "v6") else 10000)
        if v.get("role"):
            s += f" ({v['role']})"
        if v.get("current_value"):
            s += f" = {str(v['current_value'])[:30]!r}"
        for k in ("checked", "selected", "expanded"):
            if k in v:
                s += f" {k}={v[k]}"
        return s
    return v


def fields_summary(elements):
    """v5: the form's fields and their CURRENT values, first in the state, so the policy sees what is still empty
    (v4's option lists no longer repeat a dropdown's current value).  Same code in finetune/common_ft.py."""
    out = []
    for e in elements or []:
        ops, role = e.get("operations") or [], e.get("role")
        if "TYPE_TEXT" in ops or "SELECT" in ops or role == "combobox":
            v = str(e.get("value") or "").strip()
            out.append(f"{str(e.get('label', ''))[:40]} = {v[:30]!r}" if v else f"{str(e.get('label', ''))[:40]} = (empty)")
        elif role in ("checkbox", "radio", "switch") and "checked" in e:
            out.append(f"{str(e.get('label', ''))[:40]}: checked={e['checked']}")
        if len(out) >= 14:
            break
    return "; ".join(out)


def short_url(u):
    u = re.sub(r"^https?://[^/]+", "", str(u or ""))
    return u[:90] or "/"


def history_v6(history):
    """v6 history (same code in finetune/common_ft.py): last 20 actions, each with the page it led to."""
    out = []
    for h in list(history)[-20:]:
        e = {"action": str(h.get("action", ""))[:60], "kind": h.get("kind")}
        if h.get("text"):
            e["text"] = str(h["text"])[:40]
        if h.get("url") or h.get("title"):
            e["result"] = (short_url(h.get("url")) + " | " + str(h.get("title") or "")[:40]).strip(" |")
        elif h.get("page_changed") is False:
            e["result"] = "no change"
        out.append(e)
    return out


def predict(state, questions):
    """agent.predict with coarse-to-fine handling of wide choice questions.

    A choice with more than MAXOPT options is split into interleaved chunks; every chunk is a question in the same
    forward pass as the normal questions, then the chunk winners compete in a second pass.
    p(option) = p_final(winner of its chunk) * p_chunk(option)."""
    qs, plan = {}, {}
    if isinstance(state, dict) and isinstance(state.get("page"), dict) and isinstance(state["page"].get("text"), str):
        if FMT in ("v2", "v3", "v4", "v5", "v6"):   # mirror finetune/common_ft.py
            ra = state.get("recent_actions", [])
            if FMT != "v6":   # the client now also sends url/title per step; older formats were trained without them
                ra = [{k: h.get(k) for k in ("action", "kind", "text", "page_changed")} for h in ra[-10:]]
            st = {"page": {**state["page"], "text": state["page"]["text"][:{"v2": 1500, "v6": 3000}.get(FMT, 1200)]}, "recent_actions": ra}
            if FMT == "v6":
                state = {"fields": fields_summary(state.get("elements")), "recent_actions": history_v6(st["recent_actions"]), "page": st["page"]}
            else:
                state = {"fields": fields_summary(state.get("elements")), **st} if FMT == "v5" else st
        else:
            state = {**state, "page": {**state["page"], "text": state["page"]["text"][:6000]}}
    for qid, q in questions.items():
        q = dict(q)
        if isinstance(q.get("criteria"), dict):
            q["criteria"] = {k: compact(v) for k, v in q["criteria"].items()}
        keys = list(q["criteria"]) if q["type"] == "choice" and isinstance(q.get("criteria"), dict) else []
        if len(keys) <= MAXOPT:
            qs[qid] = q
            continue
        n = -(-len(keys) // MAXOPT)
        chunks = [keys[i::n] for i in range(n)]
        plan[qid] = (q, chunks)
        for ci, ch in enumerate(chunks):
            qs[f"{qid}__chunk{ci}"] = {**q, "criteria": {k: q["criteria"][k] for k in ch}}
    r = predict_fast(agent, state, qs)
    r["passes"] = 1
    if plan:
        chunk_ans = {qid: [r["answers"].pop(f"{qid}__chunk{ci}") for ci in range(len(chunks))] for qid, (q, chunks) in plan.items()}
        finals = {qid: {**q, "criteria": {a["choice"]: q["criteria"][a["choice"]] for a in chunk_ans[qid]}} for qid, (q, _) in plan.items()}
        r2 = predict_fast(agent, state, finals)
        r["passes"] = 2
        r["usage"]["input_tokens"] += r2["usage"]["input_tokens"]
        for qid, (q, chunks) in plan.items():
            fa = r2["answers"][qid]
            probs = {}
            for ca, ch in zip(chunk_ans[qid], chunks):
                pf = fa["probabilities"][ca["choice"]]
                for k in ch:
                    probs[k] = pf * ca["probabilities"][k]
            tot = sum(probs.values()) or 1.0
            probs = {k: round(v / tot, 6) for k, v in probs.items()}
            choice = max(probs, key=probs.get)
            r["answers"][qid] = {"type": "choice", "choice": choice, "probabilities": probs, "confidence": fa["confidence"],
                                 "action": fa.get("action", {}), "coarse_to_fine": {"chunks": len(chunks), "winners": [a["choice"] for a in chunk_ans[qid]]}}
    return r


class H(BaseHTTPRequestHandler):
    def log_message(self, *a): pass
    def _send(self, code, body):
        data = json.dumps(body, ensure_ascii=False).encode()
        self.send_response(code); self.send_header("Content-Type", "application/json"); self.send_header("Content-Length", str(len(data))); self.end_headers(); self.wfile.write(data)
    def do_GET(self):
        self._send(200, {"ok": True, "variant": VARIANT, "calls": len(LOG), "tau": TAU, "escalated": STATS["escalated"], "recent": LOG[-5:]})
    def do_POST(self):
        n = int(self.headers.get("Content-Length", 0)); req = json.loads(self.rfile.read(n) or b"{}")
        try:
            t = time.perf_counter()
            r = predict(req["state"], req["questions"])
            STATS["calls"] += 1
            if TAU > 0 or req.get("escalate"):
                a = r["answers"]; op = a["operation"]["choice"]; tq = op.lower() + "_target"
                conf = min(a["operation"]["confidence"], a[tq]["confidence"] if tq in a else 1.0)
                # the agent asks for System 2 itself when the fast policy is stuck or its DONE was rejected
                if (TAU > 0 and conf < TAU) or req.get("escalate"):
                    try:
                        r["answers"], esc = escalate(req["state"], req["questions"], a, req.get("escalate") or "low_confidence")
                        STATS["escalated"] += esc
                    except Exception as e:
                        print("[escalate] failed:", str(e)[:80], flush=True)
            ms = (time.perf_counter() - t) * 1000
            r["model"] = f"laya-{VARIANT}"
            nq = len(req["questions"]); nopt = sum(len(q.get("criteria") or []) for q in req["questions"].values())
            LOG.append({"ms": round(ms, 1), "questions": nq, "options": nopt, "tokens": r["usage"]["input_tokens"], "passes": r["passes"]})
            print(f"[systemone] {nq} q / {nopt} opts / {r['usage']['input_tokens']} tok / {r['passes']} pass -> {ms:.1f} ms  " +
                  ", ".join(f"{k}={v.get('choice', v.get('score', v.get('noul')))}({v['confidence']:.2f})" for k, v in r["answers"].items()), flush=True)
            self._send(200, r)
        except Exception as e:
            traceback.print_exc(); self._send(400, {"error": f"{type(e).__name__}: {e}"})

if __name__ == "__main__":
    print(f"laya systemone server on http://127.0.0.1:{PORT}/v1/systemone  ({VARIANT})", flush=True)
    ThreadingHTTPServer(("127.0.0.1", PORT), H).serve_forever()