proof-of-concept for thread-based bidi in python
diff --git a/py/requirements.txt b/py/requirements.txt index bc91501..5f706cc 100644 --- a/py/requirements.txt +++ b/py/requirements.txt
@@ -32,5 +32,6 @@ twine==4.0.2 typing_extensions==4.9.0 urllib3[socks]==2.0.7 +websocket-client==1.8.0 wsproto==1.2.0 zipp==3.17.0
diff --git a/py/selenium/webdriver/chromium/options.py b/py/selenium/webdriver/chromium/options.py index 2b305d3..26192cc 100644 --- a/py/selenium/webdriver/chromium/options.py +++ b/py/selenium/webdriver/chromium/options.py
@@ -35,6 +35,7 @@ self._extensions = [] self._experimental_options = {} self._debugger_address = None + self.add_argument("--remote-allow-origins=*") @property def binary_location(self) -> str:
diff --git a/py/selenium/webdriver/common/bidi/cdp.py b/py/selenium/webdriver/common/bidi/cdp.py index b2b3ae5..3236027 100644 --- a/py/selenium/webdriver/common/bidi/cdp.py +++ b/py/selenium/webdriver/common/bidi/cdp.py
@@ -499,3 +499,7 @@ cdp_conn = CdpConnection(ws) nursery.start_soon(cdp_conn._reader_task) return cdp_conn + +from selenium.webdriver.common.bidi.websocket_connection import WebSocketConnection +def connect_cdp_sync(url) -> WebSocketConnection: + return WebSocketConnection(url)
diff --git a/py/selenium/webdriver/common/bidi/websocket_connection.py b/py/selenium/webdriver/common/bidi/websocket_connection.py new file mode 100644 index 0000000..410453f --- /dev/null +++ b/py/selenium/webdriver/common/bidi/websocket_connection.py
@@ -0,0 +1,121 @@ +import json +import socket +import ssl +import time +import threading +import logging +from urllib.parse import urlparse + +from websocket import WebSocketApp, WebSocketException, enableTrace +# enableTrace(True) + +logger = logging.getLogger("websocket") + +class WebSocketConnection: + # CONNECTION_ERRORS = [ + # ConnectionResetError, # connection is aborted (browser process was killed) + # BrokenPipeError # broken pipe (browser process was killed) + # ] + + RESPONSE_WAIT_TIMEOUT = 2 # TODO 30 + RESPONSE_WAIT_INTERVAL = 0.1 + + MAX_LOG_MESSAGE_SIZE = 9999 + + def __init__(self, url): + self.callback_threads = [] + self.callbacks = {} + self.messages = {} + self.id = 0 + self.started = False + + self.session_id = None + self.url = url + + self.socket_thread = self.attach_socket_listener() + self.wait_until(lambda: self.started) + + def close(self): + for thread in self.callback_threads: + thread.join() + self.socket_thread.join() + self.ws.close() + + def send_cmd(self, cmd): + id = self.next_id() + payload = next(cmd) + payload['id'] = id + if self.session_id: + payload["sessionId"] = self.session_id + data = json.dumps(payload) + logger.warning(f"WebSocket -> {data}"[:self.MAX_LOG_MESSAGE_SIZE]) + self.ws.send(data) + return self.wait_until(lambda: self.retrieve_message(id, cmd)) + + def retrieve_message(self, id, cmd): + if id in self.messages: + message = self.messages.pop(id) + try: + _ = cmd.send(message["result"]) + raise InternalError("The command's generator function did not exit when expected!") + except StopIteration as exit: + return exit.value + else: + return None + + def attach_socket_listener(self): + def on_open(ws): + self.started = True + + def on_message(ws, message): + logger.warning(message) + message = self.process_frame(message) + + if 'method' in message: + params = message['params'] + for callback in self.callbacks.get(message['method'], []): + callback(params) + + def on_error(ws, error): + logger.warning(f"WebSocket error: {error}") + ws.close() + + def run_socket(): + # TODO: Support wss + # self.ws.run_forever(sslopt={"cert_reqs": ssl.CERT_NONE}) + self.ws.run_forever() + + self.ws = WebSocketApp(self.url, on_message=on_message, on_error=on_error) + thread = threading.Thread(target=run_socket) + thread.start() + return thread + + def process_frame(self, frame): + message = frame + + # Firefox will periodically fail on unparsable empty frame + if not message: + return {} + + message = json.loads(message) + logger.warning(f"WebSocket <- {message}"[:self.MAX_LOG_MESSAGE_SIZE]) + + if 'id' in message: + self.messages[message['id']] = message + + return message + + def wait_until(self, condition): + timeout = self.RESPONSE_WAIT_TIMEOUT + interval = self.RESPONSE_WAIT_INTERVAL + + while timeout > 0: + result = condition() + if result: + return result + timeout -= interval + time.sleep(interval) + + def next_id(self): + self.id += 1 + return self.id
diff --git a/py/selenium/webdriver/remote/webdriver.py b/py/selenium/webdriver/remote/webdriver.py index c1fa511..74977f2 100644 --- a/py/selenium/webdriver/remote/webdriver.py +++ b/py/selenium/webdriver/remote/webdriver.py
@@ -1017,6 +1017,33 @@ """ return self.execute(Command.GET_LOG, {"type": log_type})["value"] + @property + def bidi(self): + global cdp + import_cdp() + if self.caps.get("se:cdp"): + ws_url = self.caps.get("se:cdp") + version = self.caps.get("se:cdpVersion").split(".")[0] + else: + version, ws_url = self._get_cdp_details() + + if not ws_url: + raise WebDriverException("Unable to find url to connect to from capabilities") + + global devtools + devtools = cdp.import_devtools(version) + conn = cdp.connect_cdp_sync(ws_url) + targets = conn.send_cmd(devtools.target.get_targets()) + target_id = targets[0].target_id + session = conn.send_cmd(devtools.target.attach_to_target(target_id, True)) + conn.session_id = session + return conn + + def on_log_event(self, type, callback): + bidi = self.bidi + bidi.send_cmd(devtools.runtime.enable()) + bidi.callbacks['Runtime.consoleAPICalled'] = [callback] + @asynccontextmanager async def bidi_connection(self): global cdp
diff --git a/py/test/selenium/webdriver/common/bidi_tests.py b/py/test/selenium/webdriver/common/bidi_tests.py index fe09212..fa4b107 100644 --- a/py/test/selenium/webdriver/common/bidi_tests.py +++ b/py/test/selenium/webdriver/common/bidi_tests.py
@@ -21,60 +21,76 @@ from selenium.webdriver.support import expected_conditions as EC from selenium.webdriver.support.ui import WebDriverWait +def test_loads(driver, pages): + events = [] + driver.on_log_event("console", lambda event: events.append(event)) + driver.execute_script("console.log('I love cheese')") + WebDriverWait(driver, 5).until(lambda _: events) -@pytest.mark.xfail_firefox(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") -@pytest.mark.xfail_remote(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") -@pytest.mark.xfail_safari -async def test_check_console_messages(driver, pages): - async with driver.bidi_connection() as session: - log = Log(driver, session) - pages.load("javascriptPage.html") - from selenium.webdriver.common.bidi.console import Console - - async with log.add_listener(Console.LOG) as messages: - driver.execute_script("console.log('I love cheese')") - assert messages["message"] == "I love cheese" + assert events[0]['args'][0]['value'] == "I love cheese" -@pytest.mark.xfail_firefox(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") -@pytest.mark.xfail_remote(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") -@pytest.mark.xfail_safari -async def test_check_error_console_messages(driver, pages): - async with driver.bidi_connection() as session: - log = Log(driver, session) - pages.load("javascriptPage.html") - from selenium.webdriver.common.bidi.console import Console - async with log.add_listener(Console.ERROR) as messages: - driver.execute_script('console.error("I don\'t cheese")') - driver.execute_script("console.log('I love cheese')") - assert messages["message"] == "I don't cheese" +# driver.on_log_event("console", lambda x: print(x)) + +# def log(event): +# print(event) +# driver.on_log_event("console", log) -@pytest.mark.xfail_firefox -@pytest.mark.xfail_safari -@pytest.mark.xfail_remote -async def test_collect_js_exceptions(driver, pages): - async with driver.bidi_connection() as session: - log = Log(driver, session) - pages.load("javascriptPage.html") - async with log.add_js_error_listener() as exceptions: - driver.find_element(By.ID, "throwing-mouseover").click() - assert exceptions is not None - assert exceptions.exception_details.stack_trace.call_frames[0].function_name == "onmouseover" +# @pytest.mark.xfail_firefox(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") +# @pytest.mark.xfail_remote(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") +# @pytest.mark.xfail_safari +# async def test_check_console_messages(driver, pages): +# async with driver.bidi_connection() as session: +# log = Log(driver, session) +# pages.load("javascriptPage.html") +# from selenium.webdriver.common.bidi.console import Console + +# async with log.add_listener(Console.LOG) as messages: +# driver.execute_script("console.log('I love cheese')") +# assert messages["message"] == "I love cheese" -@pytest.mark.xfail_firefox -@pytest.mark.xfail_safari -@pytest.mark.xfail_remote -async def test_collect_log_mutations(driver, pages): - async with driver.bidi_connection() as session: - log = Log(driver, session) - async with log.mutation_events() as event: - pages.load("dynamic.html") - driver.find_element(By.ID, "reveal").click() - WebDriverWait(driver, 5).until(EC.visibility_of(driver.find_element(By.ID, "revealed"))) +# @pytest.mark.xfail_firefox(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") +# @pytest.mark.xfail_remote(reason="https://bugzilla.mozilla.org/show_bug.cgi?id=1819965") +# @pytest.mark.xfail_safari +# async def test_check_error_console_messages(driver, pages): +# async with driver.bidi_connection() as session: +# log = Log(driver, session) +# pages.load("javascriptPage.html") +# from selenium.webdriver.common.bidi.console import Console - assert event["attribute_name"] == "style" - assert event["current_value"] == "" - assert event["old_value"] == "display:none;" +# async with log.add_listener(Console.ERROR) as messages: +# driver.execute_script('console.error("I don\'t cheese")') +# driver.execute_script("console.log('I love cheese')") +# assert messages["message"] == "I don't cheese" + + +# @pytest.mark.xfail_firefox +# @pytest.mark.xfail_safari +# @pytest.mark.xfail_remote +# async def test_collect_js_exceptions(driver, pages): +# async with driver.bidi_connection() as session: +# log = Log(driver, session) +# pages.load("javascriptPage.html") +# async with log.add_js_error_listener() as exceptions: +# driver.find_element(By.ID, "throwing-mouseover").click() +# assert exceptions is not None +# assert exceptions.exception_details.stack_trace.call_frames[0].function_name == "onmouseover" + + +# @pytest.mark.xfail_firefox +# @pytest.mark.xfail_safari +# @pytest.mark.xfail_remote +# async def test_collect_log_mutations(driver, pages): +# async with driver.bidi_connection() as session: +# log = Log(driver, session) +# async with log.mutation_events() as event: +# pages.load("dynamic.html") +# driver.find_element(By.ID, "reveal").click() +# WebDriverWait(driver, 5).until(EC.visibility_of(driver.find_element(By.ID, "revealed"))) + +# assert event["attribute_name"] == "style" +# assert event["current_value"] == "" +# assert event["old_value"] == "display:none;"