aroniscunt/Browser_Web_UI_Automation
0
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 