From 20ad25989d81ebcaf0f43133e68ae1bc00273eb5 Mon Sep 17 00:00:00 2001 From: alckasoc Date: Tue, 14 Jan 2025 18:32:18 -0800 Subject: [PATCH] add uielement to osworldaci; rename osworldaci; update readme --- OSWorld.md | 4 +- README.md | 10 +-- .../aci/{OSWorldACI.py => LinuxOSACI.py} | 87 +++++++++++++++++-- gui_agents/cli_app.py | 7 +- 4 files changed, 92 insertions(+), 16 deletions(-) rename gui_agents/aci/{OSWorldACI.py => LinuxOSACI.py} (93%) diff --git a/OSWorld.md b/OSWorld.md index d7c8497..615b2aa 100644 --- a/OSWorld.md +++ b/OSWorld.md @@ -26,7 +26,7 @@ We suggest creating a separate conda environment for each repository to avoid de After completing the setup instructions, import the GraphSearchAgent into the run.py file in OSWorld. The GraphSearchAgent is the parent agent used in the Agent S framework. To understand the architecture of this GraphSearchAgent, refer to [Agent S Architecture](images/agent_s_architecture.pdf). ``` -from gui_agents.aci.OSWorldACI import OSWorldACI +from gui_agents.aci.LinuxOSACI import LinuxACI from gui_agents.core.AgentS import GraphSearchAgent ``` @@ -49,7 +49,7 @@ engine_params = { "model": args.model, } -grounding_agent = OSWorldACI(vm_version=args.vm_version) +grounding_agent = LinuxACI(vm_version=args.vm_version) agent = GraphSearchAgent( engine_params, grounding_agent, diff --git a/README.md b/README.md index 3751123..cdf83aa 100644 --- a/README.md +++ b/README.md @@ -138,7 +138,6 @@ import io from gui_agents.core.AgentS import GraphSearchAgent import platform - if platform.system() == "Darwin": from gui_agents.aci.MacOSACI import MacOSACI, UIElement grounding_agent = MacOSACI() @@ -146,8 +145,8 @@ elif platform.system() == "Windows": from gui_agents.aci.WindowsOSACI import WindowsACI, UIElement grounding_agent = WindowsACI() elif platform.system() == "Linux": - from gui_agents.aci.OSWorldACI import OSWorldACI, get_acc_tree - grounding_agent = OSWorldACI() + from gui_agents.aci.LinuxOSACI import LinuxACI, UIElement + grounding_agent = LinuxACI() else: raise ValueError("Unsupported platform") @@ -172,10 +171,7 @@ screenshot.save(buffered, format="PNG") screenshot_bytes = buffered.getvalue() # Get accessibility tree. -if platform.system() != "Linux": - acc_tree = UIElement.systemWideElement() -elif platform.system() == "Linux": - acc_tree = get_acc_tree() +acc_tree = UIElement.systemWideElement() obs = { "screenshot": screenshot_bytes, diff --git a/gui_agents/aci/OSWorldACI.py b/gui_agents/aci/LinuxOSACI.py similarity index 93% rename from gui_agents/aci/OSWorldACI.py rename to gui_agents/aci/LinuxOSACI.py index 9aa65e7..9c36411 100644 --- a/gui_agents/aci/OSWorldACI.py +++ b/gui_agents/aci/LinuxOSACI.py @@ -3,7 +3,7 @@ import logging import os import time import xml.etree.ElementTree as ET -from typing import Dict, List, Optional, Tuple +from typing import Dict, List, Optional, Tuple, Any, Sequence import requests import torch @@ -25,11 +25,7 @@ if platform.system() == "Linux": from lxml.etree import _Element from typing import Optional, Dict, Any, List - import platform - from typing import Any, Optional, Sequence import lxml.etree - from flask import jsonify - from lxml.etree import _Element import concurrent.futures logger = logging.getLogger("desktopenv.agent") @@ -75,7 +71,7 @@ def agent_action(func): return func -class OSWorldACI(ACI): +class LinuxACI(ACI): def __init__(self, top_app=None, vm_version="new", top_app_only=True, ocr=True): self.active_apps = set() self.top_app = top_app @@ -879,3 +875,82 @@ def get_acc_tree() -> str: xml_node.append(xml_tree) acc_tree = lxml.etree.tostring(xml_node, encoding="unicode") return acc_tree + + +class UIElement(object): + def __init__(self, node: Accessible): + self.node: Accessible = node + + def getAttributeNames(self): + attributes = self.node.getAttributes() + + @staticmethod + def systemWideElement(): + # desktop = pyatspi.Registry.getDesktop(0) + # for app in desktop: + # for window in app: + # if window.getState().contains(pyatspi.STATE_ACTIVE): + # active_node = app + # return UIElement(active_node) + return get_acc_tree() + + @property + def states(self): + state_names = [] + states: List[StateType] = self.node.getState().get_states() + for st in states: + state_name: str = StateType._enum_lookup[st] + state_names.append(state_name) + return state_names + + @property + def attributes(self): + try: + attributes: List[str] = self.node.getAttributes() + attribute_dict = {} + for attrbt in attributes: + attribute_name: str + attribute_value: str + attribute_name, attribute_value = attrbt.split(":", maxsplit=1) + attribute_dict[attribute_name] = attribute_value + return attribute_dict + except NotImplementedError: + return None + + @property + def component(self): + try: + component: Component = self.node.queryComponent() + return component + except NotImplementedError: + return None + + @property + def value(self): + try: + value: ATValue = self.node.queryValue() + return value + except NotImplementedError: + return None + + @property + def text(self): + try: + text_obj: ATText = self.node.queryText() + except NotImplementedError: + return "" + else: + text: str = text_obj.getText(0, text_obj.characterCount) + text = text.replace("\ufffc", "").replace("\ufffd", "") + return text + + @property + def role(self): + return self.node.getRoleName() + + def children(self): + """Return list of children of the current node""" + return list(self.node) + + def __repr__(self): + return "UIElement%s" % (self.node) diff --git a/gui_agents/cli_app.py b/gui_agents/cli_app.py index ce0e1a5..cfe8587 100644 --- a/gui_agents/cli_app.py +++ b/gui_agents/cli_app.py @@ -14,7 +14,10 @@ if platform.system() == "Darwin": from gui_agents.aci.MacOSACI import MacOSACI, UIElement elif platform.system() == "Windows": current_platform = "windows" - from gui_agents.aci.WindowsOSACI import UIElement, WindowsACI + from gui_agents.aci.WindowsOSACI import WindowsACI, UIElement +elif platform.system() == "Linux": + current_platform = "ubuntu" + from gui_agents.aci.LinuxOSACI import LinuxACI, UIElement else: raise ValueError("Unsupported platform") @@ -155,6 +158,8 @@ def main(): grounding_agent = MacOSACI() elif platform.system() == "Windows": grounding_agent = WindowsACI() + elif platform.system() == "Linux": + grounding_agent = LinuxACI() else: raise ValueError("Unsupported platform")