DeskForge / web_agent.py
Saidgurbuz's picture
Catch agent loops: compare screens by thumbnail, count repeated actions; fix back on in-page navigation; simpler Docling task
25f9914 verified
Raw History Blame Contribute Delete
15.3 kB
"""The web agent: the paper's planner + DeskForge grounder, driving a real browser one step at a time.
The page drives the loop through three server functions so it can animate between
steps: `web_plan` (screenshot -> planner -> DeskForge, one GPU call), then
`web_act` (perform the action in the browser and return the new screenshot).
Runs are kept here between calls and closed when finished, stopped, or idle.
"""
import base64
import hashlib
import io
import json
import threading
import time
import traceback
import uuid
from dataclasses import dataclass, field
from PIL import Image
from web_browser import VIEWPORT, BrowserHost, ensure_chromium
from web_plan import planner_user_text, typing_blocked
MAX_STEPS = 30
MAX_RUNS = 4 # browsers working at the same time; later visitors wait for a free one
IDLE_SECONDS = 180 # a run nobody polls for this long is closed (tab closed, network lost; a busy queue is not idle)
RUN_SECONDS = 8 * 60
MAX_TASK_CHARS = 300
SPARE_SECONDS = 10 * 60 # a pre-opened start page older than this is replaced
def _thumbnail(png: bytes) -> bytes:
return Image.open(io.BytesIO(png)).convert("L").resize((48, 30), Image.BILINEAR).tobytes()
def _close(a: bytes, b: bytes) -> bool:
"""Two thumbnails show the same screen (mean difference under ~1% of the gray range)."""
return len(a) == len(b) and sum(abs(x - y) for x, y in zip(a, b)) / len(a) < 2.5
def _jpeg(png: bytes, quality=80) -> str:
image = Image.open(io.BytesIO(png)).convert("RGB")
buf = io.BytesIO()
image.save(buf, format="JPEG", quality=quality, optimize=True)
return "data:image/jpeg;base64," + base64.b64encode(buf.getvalue()).decode()
@dataclass
class Run:
id: str
task: str
client: str
session: object
created: float = field(default_factory=time.time)
seen: float = field(default_factory=time.time)
step: int = 0
memory: str = ""
history: list = field(default_factory=list)
pending: dict | None = None # the decision web_act will perform
shot: bytes = b"" # the screenshot the next decision is made on
shown: str = "" # hash of the last screenshot sent to the page
errors: int = 0 # consecutive failed steps
last_sig: bytes | None = None # stuck detection: the last screen's thumbnail, and for how many steps it has not changed
same: int = 0
done: bool = False
lock: threading.Lock = field(default_factory=threading.Lock)
class WebAgent:
def __init__(self, decide, display=None, max_runs=MAX_RUNS, max_steps=MAX_STEPS):
"""decide(user_text, PIL image) -> dict from web_plan.decide_with (points as screen fractions)."""
self.decide = decide
ensure_chromium()
self.host = BrowserHost(display)
threading.Thread(target=self._warm, daemon=True, name="web-warm").start()
self.max_runs, self.max_steps = max_runs, max_steps
self.runs: dict[str, Run] = {}
self._lock = threading.Lock()
self._spare = None
threading.Thread(target=self._reaper, daemon=True, name="web-reaper").start()
# ---- lifecycle --------------------------------------------------------------
def _warm(self):
"""Start Chromium in the background and open a spare tab, so the first visitor does not wait."""
try:
self.host.call(self.host._ensure(), timeout=180)
except Exception:
traceback.print_exc()
try:
ok = self.host.call(self.host.family_dns_ok(), timeout=15)
except Exception:
ok = False
print(f"[web] family DNS filter {'on' if ok else 'UNREACHABLE: pages are not filtered for adult/malware sites'}",
flush=True)
self._refill()
def _refill(self):
"""Keep one private context already on the start page; a new run takes it instead of waiting."""
try:
session = self.host.call(self.host.new_session(), timeout=120)
except Exception:
traceback.print_exc()
return
session.ready_at = time.time()
with self._lock:
spare, self._spare = self._spare, session
if spare:
self.host.call(spare.close(), timeout=20)
def _take_spare(self):
with self._lock:
session, self._spare = self._spare, None
threading.Thread(target=self._refill, daemon=True, name="web-refill").start()
if session and time.time() - session.ready_at > SPARE_SECONDS:
self.host.call(session.close(), timeout=20) # a start page left open too long may be stale
session = None
return session
def _reaper(self):
while True:
time.sleep(5)
now = time.time()
for run in list(self.runs.values()):
if now - run.seen > IDLE_SECONDS or now - run.created > RUN_SECONDS + 60:
self._close(run.id)
def _close(self, run_id):
with self._lock:
run = self.runs.pop(run_id, None)
if run:
print(f"[web] close {run_id} after {run.step} steps", flush=True)
if run:
run.done = True
try:
self.host.call(run.session.close(), timeout=20)
except Exception:
pass
def _get(self, run_id) -> Run:
run = self.runs.get(run_id)
if run is None:
print(f"[web] unknown run {run_id!r}; active: {list(self.runs)}", flush=True)
raise KeyError("This run has ended. Start a new one.")
run.seen = time.time()
return run
def _image(self, run, png):
"""The screenshot as a JPEG data URI, or None when the page already shows it."""
digest = hashlib.md5(png).hexdigest()
if digest == run.shown:
return None
run.shown = digest
return _jpeg(png)
def busy(self):
return len(self.runs) >= self.max_runs
# ---- the three steps the page calls ----------------------------------------------
def start(self, task: str, client: str) -> dict:
task = " ".join((task or "").split())[:MAX_TASK_CHARS]
if len(task) < 4:
return {"error": "Write a task first, e.g. “Find the weather in Zurich this weekend”."}
for old in [r.id for r in self.runs.values() if r.client == client]:
self._close(old) # one run per browser tab
with self._lock:
if len(self.runs) >= self.max_runs:
return {"busy": True, "active": len(self.runs)}
run_id = uuid.uuid4().hex[:12]
self.runs[run_id] = run = Run(run_id, task, client, session=None)
try:
run.session = self._take_spare() or self.host.call(self.host.new_session(), timeout=120)
run.shot = self.host.call(run.session.screenshot(), timeout=30)
info = self.host.call(run.session.info(), timeout=10)
except Exception as exc:
traceback.print_exc()
self._close(run_id)
return {"error": f"The browser could not start ({type(exc).__name__}). Please try again."}
return {"run": run_id, "task": task, "max_steps": self.max_steps, "image": self._image(run, run.shot),
"viewport": [VIEWPORT["width"], VIEWPORT["height"]], **info}
def plan(self, run_id: str) -> dict:
run = self._get(run_id)
with run.lock:
if run.done:
return {"done": True}
if run.step >= self.max_steps or time.time() - run.created > RUN_SECONDS:
return self._finish(run, "", f"Stopped after {run.step} steps (the demo's limit).", limit=True)
try:
info = self.host.call(run.session.info(), timeout=10)
fresh = self.host.call(run.session.screenshot(), timeout=30)
except Exception:
traceback.print_exc()
run.errors += 1
if run.errors >= 3:
return self._finish(run, "", "Stopped: the browser stopped responding.")
try: # a page that stopped responding is replaced by a fresh tab on the start page
self.host.call(run.session.recover(), timeout=40)
except Exception:
traceback.print_exc()
run.history.append({"action": {"note": "the page stopped responding"},
"execution_error": "the tab was reopened on the start page"})
return {"step": run.step, "error": "The page stopped responding, so the agent reopened the start page.",
"t_plan": 0, "t_ground": 0, "t_total": 0}
run.shot = fresh
stuck = self._check_stuck(run, fresh)
if stuck:
return self._finish(run, "", stuck)
text = planner_user_text(run.task, run.step + 1, self.max_steps, run.memory, run.history,
info["url"], info["title"])
image = Image.open(io.BytesIO(fresh)).convert("RGB")
t0 = time.perf_counter()
try:
d = self.decide(text, image)
except Exception as exc: # e.g. the visitor's ZeroGPU quota ran out
return {**self._finish(run, "", str(exc) or type(exc).__name__), "fatal": True}
wall = time.perf_counter() - t0
run.step += 1
out = {"step": run.step, "image": self._image(run, fresh), "url": info["url"], "title": info["title"],
"t_plan": round(d.get("t_plan", 0), 2), "t_ground": round(d.get("t_ground", 0), 2),
"t_total": round(wall, 2)}
if d.get("error"):
run.errors += 1
plan = d.get("plan") or {}
run.history.append({"action": {k: v for k, v in plan.items() if k != "memory"} or {"raw": d["raw"][:200]},
"execution_error": d["error"]})
run.pending = None
out.update(error=d["error"], plan=plan or None)
if run.errors >= 3:
out.update(self._finish(run, "", "Stopped: the planner kept producing actions it could not run."))
return out
plan = d["plan"]
if plan.get("memory"):
run.memory = plan["memory"]
points = [d["points"][k] for k in ("target", "destination") if k in d["points"]]
run.pending = {"plan": plan, "points": points}
out.update(plan=plan, points=points)
return out
def act(self, run_id: str) -> dict:
run = self._get(run_id)
with run.lock:
if run.done:
return {"done": True}
pending, run.pending = run.pending, None
if pending is None:
return {"skipped": True}
plan, points = pending["plan"], pending["points"]
if plan["action"] == "finish":
return self._finish(run, plan.get("answer", ""), "")
error = typing_blocked(plan["text"]) if plan["action"] == "type" else ""
if not error:
px = [(round(x * (VIEWPORT["width"] - 1)), round(y * (VIEWPORT["height"] - 1))) for x, y in points]
try:
error = self.host.call(run.session.act(plan, px), timeout=45)
except Exception as exc:
error = f"the action failed ({type(exc).__name__})"
blocked = run.session.blocked
run.history.append({"action": {k: v for k, v in plan.items() if k != "memory"},
"execution_error": error or (f"blocked: {blocked}" if blocked else None)})
run.errors = run.errors + 1 if error else 0
try:
run.shot = self.host.call(run.session.screenshot(), timeout=30)
info = self.host.call(run.session.info(), timeout=10)
except Exception:
traceback.print_exc()
return {"image": None, "error": error, "blocked": blocked}
return {"image": self._image(run, run.shot), "error": error, "blocked": blocked, **info}
def _check_stuck(self, run, png):
"""Tell the planner when it is going in circles, and give up when it keeps doing so.
"Unchanged" compares small grayscale thumbnails, so clocks, cursors and animations do not hide a loop;
repeating the very same action counts too, whatever the screen does.
"""
sig = _thumbnail(png)
same = run.last_sig is not None and _close(sig, run.last_sig)
run.same = run.same + 1 if same else 0
run.last_sig = sig
acts = [json.dumps(h["action"], sort_keys=True) for h in run.history if "action" in h.get("action", {})]
repeat = 0
for a in reversed(acts):
if a != acts[-1]:
break
repeat += 1
if run.same >= 6:
return "Stopped: the page did not change for six steps."
if repeat >= 8:
return "Stopped: the agent kept repeating the same action."
if run.same >= 3 or repeat >= 4:
what = f"the page has not changed for {run.same} steps" if run.same >= 3 else f"the same action was repeated {repeat} times"
run.history.append({"action": {"note": what},
"execution_error": "it is not working: do something different (another element, back, "
"another site) or finish and say what blocked you"})
return ""
def _finish(self, run, answer, note, limit=False):
run.done = True
out = {"done": True, "answer": answer, "note": note, "limit": limit, "steps": run.step,
"seconds": round(time.time() - run.created, 1)}
threading.Thread(target=self._close, args=(run.id,), daemon=True).start()
return out
def stop(self, run_id: str) -> dict:
self._close(run_id)
return {"stopped": True}
# ---- server functions for the gr.HTML view (arguments arrive as one list) ----------
def server_functions(self):
agent = self
def first(args): # one JS argument arrives bare, several arrive as a list
return str(args[0] if isinstance(args, (list, tuple)) and args else args or "")
def web_start(args: list) -> dict:
task, client = (list(args) + ["", ""])[:2] if isinstance(args, (list, tuple)) else (args, "")
return agent.start(str(task), str(client))
def web_plan(args: list) -> dict:
try:
return agent.plan(first(args))
except KeyError as exc:
return {"done": True, "note": str(exc.args[0])}
def web_act(args: list) -> dict:
try:
return agent.act(first(args))
except KeyError as exc:
return {"done": True, "note": str(exc.args[0])}
def web_stop(args: list) -> dict:
return agent.stop(first(args))
def web_status(args: list) -> dict:
return {"busy": agent.busy(), "active": len(agent.runs), "max": agent.max_runs}
return [web_start, web_plan, web_act, web_stop, web_status]