CoolFace
Apppublic

aroniscunt/Browser_Web_UI_Automation

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
custom_controller.py183 linesDownload Raw Back to controller
1import pdb
2
3import pyperclip
4from typing import Optional, Type, Callable, Dict, Any, Union, Awaitable, TypeVar
5from pydantic import BaseModel
6from browser_use.agent.views import ActionResult
7from browser_use.browser.context import BrowserContext
8from browser_use.controller.service import Controller, DoneAction
9from browser_use.controller.registry.service import Registry, RegisteredAction
10from main_content_extractor import MainContentExtractor
11from browser_use.controller.views import (
12    ClickElementAction,
13    DoneAction,
14    ExtractPageContentAction,
15    GoToUrlAction,
16    InputTextAction,
17    OpenTabAction,
18    ScrollAction,
19    SearchGoogleAction,
20    SendKeysAction,
21    SwitchTabAction,
22)
23import logging
24import inspect
25import asyncio
26import os
27from langchain_core.language_models.chat_models import BaseChatModel
28from browser_use.agent.views import ActionModel, ActionResult
29
30from src.utils.mcp_client import create_tool_param_model, setup_mcp_client_and_tools
31
32from browser_use.utils import time_execution_sync
33
34logger = logging.getLogger(__name__)
35
36Context = TypeVar('Context')
37
38
39class CustomController(Controller):
40    def __init__(self, exclude_actions: list[str] = [],
41                 output_model: Optional[Type[BaseModel]] = None,
42                 ask_assistant_callback: Optional[Union[Callable[[str, BrowserContext], Dict[str, Any]], Callable[
43                     [str, BrowserContext], Awaitable[Dict[str, Any]]]]] = None,
44                 ):
45        super().__init__(exclude_actions=exclude_actions, output_model=output_model)
46        self._register_custom_actions()
47        self.ask_assistant_callback = ask_assistant_callback
48        self.mcp_client = None
49        self.mcp_server_config = None
50
51    def _register_custom_actions(self):
52        """Register all custom browser actions"""
53
54        @self.registry.action(
55            "When executing tasks, prioritize autonomous completion. However, if you encounter a definitive blocker "
56            "that prevents you from proceeding independently – such as needing credentials you don't possess, "
57            "requiring subjective human judgment, needing a physical action performed, encountering complex CAPTCHAs, "
58            "or facing limitations in your capabilities – you must request human assistance."
59        )
60        async def ask_for_assistant(query: str, browser: BrowserContext):
61            if self.ask_assistant_callback:
62                if inspect.iscoroutinefunction(self.ask_assistant_callback):
63                    user_response = await self.ask_assistant_callback(query, browser)
64                else:
65                    user_response = self.ask_assistant_callback(query, browser)
66                msg = f"AI ask: {query}. User response: {user_response['response']}"
67                logger.info(msg)
68                return ActionResult(extracted_content=msg, include_in_memory=True)
69            else:
70                return ActionResult(extracted_content="Human cannot help you. Please try another way.",
71                                    include_in_memory=True)
72
73        @self.registry.action(
74            'Upload file to interactive element with file path ',
75        )
76        async def upload_file(index: int, path: str, browser: BrowserContext, available_file_paths: list[str]):
77            if path not in available_file_paths:
78                return ActionResult(error=f'File path {path} is not available')
79
80            if not os.path.exists(path):
81                return ActionResult(error=f'File {path} does not exist')
82
83            dom_el = await browser.get_dom_element_by_index(index)
84
85            file_upload_dom_el = dom_el.get_file_upload_element()
86
87            if file_upload_dom_el is None:
88                msg = f'No file upload element found at index {index}'
89                logger.info(msg)
90                return ActionResult(error=msg)
91
92            file_upload_el = await browser.get_locate_element(file_upload_dom_el)
93
94            if file_upload_el is None:
95                msg = f'No file upload element found at index {index}'
96                logger.info(msg)
97                return ActionResult(error=msg)
98
99            try:
100                await file_upload_el.set_input_files(path)
101                msg = f'Successfully uploaded file to index {index}'
102                logger.info(msg)
103                return ActionResult(extracted_content=msg, include_in_memory=True)
104            except Exception as e:
105                msg = f'Failed to upload file to index {index}: {str(e)}'
106                logger.info(msg)
107                return ActionResult(error=msg)
108
109    @time_execution_sync('--act')
110    async def act(
111            self,
112            action: ActionModel,
113            browser_context: Optional[BrowserContext] = None,
114            #
115            page_extraction_llm: Optional[BaseChatModel] = None,
116            sensitive_data: Optional[Dict[str, str]] = None,
117            available_file_paths: Optional[list[str]] = None,
118            #
119            context: Context | None = None,
120    ) -> ActionResult:
121        """Execute an action"""
122
123        try:
124            for action_name, params in action.model_dump(exclude_unset=True).items():
125                if params is not None:
126                    if action_name.startswith("mcp"):
127                        # this is a mcp tool
128                        logger.debug(f"Invoke MCP tool: {action_name}")
129                        mcp_tool = self.registry.registry.actions.get(action_name).function
130                        result = await mcp_tool.ainvoke(params)
131                    else:
132                        result = await self.registry.execute_action(
133                            action_name,
134                            params,
135                            browser=browser_context,
136                            page_extraction_llm=page_extraction_llm,
137                            sensitive_data=sensitive_data,
138                            available_file_paths=available_file_paths,
139                            context=context,
140                        )
141
142                    if isinstance(result, str):
143                        return ActionResult(extracted_content=result)
144                    elif isinstance(result, ActionResult):
145                        return result
146                    elif result is None:
147                        return ActionResult()
148                    else:
149                        raise ValueError(f'Invalid action result type: {type(result)} of {result}')
150            return ActionResult()
151        except Exception as e:
152            raise e
153
154    async def setup_mcp_client(self, mcp_server_config: Optional[Dict[str, Any]] = None):
155        self.mcp_server_config = mcp_server_config
156        if self.mcp_server_config:
157            self.mcp_client = await setup_mcp_client_and_tools(self.mcp_server_config)
158            self.register_mcp_tools()
159
160    def register_mcp_tools(self):
161        """
162        Register the MCP tools used by this controller.
163        """
164        if self.mcp_client:
165            for server_name in self.mcp_client.server_name_to_tools:
166                for tool in self.mcp_client.server_name_to_tools[server_name]:
167                    tool_name = f"mcp.{server_name}.{tool.name}"
168                    self.registry.registry.actions[tool_name] = RegisteredAction(
169                        name=tool_name,
170                        description=tool.description,
171                        function=tool,
172                        param_model=create_tool_param_model(tool),
173                    )
174                    logger.info(f"Add mcp tool: {tool_name}")
175                logger.debug(
176                    f"Registered {len(self.mcp_client.server_name_to_tools[server_name])} mcp tools for {server_name}")
177        else:
178            logger.warning(f"MCP client not started.")
179
180    async def close_mcp_client(self):
181        if self.mcp_client:
182            await self.mcp_client.__aexit__(None, None, None)
183