Skip to content
Merged
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 60 additions & 38 deletions py/selenium/webdriver/common/bidi/network.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,19 +15,23 @@
# specific language governing permissions and limitations
# under the License.

from collections.abc import Callable
from typing import Any

from selenium.webdriver.common.bidi.common import command_builder
from selenium.webdriver.remote.websocket_connection import WebSocketConnection


class NetworkEvent:
"""Represents a network event."""

def __init__(self, event_class, **kwargs):
def __init__(self, event_class: str, **kwargs: Any) -> None:
self.event_class = event_class
self.params = kwargs

@classmethod
def from_json(cls, json):
return cls(event_class=json.get("event_class"), **json)
def from_json(cls, json: dict[str, Any]) -> "NetworkEvent":
Comment thread
cgoldberg marked this conversation as resolved.
Outdated
return cls(event_class=json.get("event_class", ""), **json)


class Network:
Expand All @@ -47,13 +51,18 @@ class Network:
"auth_required": "authRequired",
}

def __init__(self, conn):
def __init__(self, conn: WebSocketConnection) -> None:
self.conn = conn
self.intercepts = []
self.callbacks = {}
self.subscriptions = {}
self.intercepts: list[str] = []
self.callbacks: dict[str | int, Any] = {}
self.subscriptions: dict[str, list[int]] = {}

def _add_intercept(self, phases=None, contexts=None, url_patterns=None):
def _add_intercept(
self,
phases: list[str] | None = None,
contexts: list[str] | None = None,
url_patterns: list[Any] | None = None,
) -> dict[str, Any]:
"""Add an intercept to the network.

Args:
Expand All @@ -77,11 +86,11 @@ def _add_intercept(self, phases=None, contexts=None, url_patterns=None):
params["phases"] = ["beforeRequestSent"]
cmd = command_builder("network.addIntercept", params)

result = self.conn.execute(cmd)
result: dict[str, Any] = self.conn.execute(cmd)
self.intercepts.append(result["intercept"])
return result

def _remove_intercept(self, intercept=None):
def _remove_intercept(self, intercept: str | None = None) -> None:
"""Remove a specific intercept, or all intercepts.

Args:
Expand All @@ -105,7 +114,7 @@ def _remove_intercept(self, intercept=None):
except Exception as e:
raise Exception(f"Exception: {e}")

def _on_request(self, event_name, callback):
def _on_request(self, event_name: str, callback: Callable[["Request"], Any]) -> int:
Comment thread
cgoldberg marked this conversation as resolved.
Outdated
"""Set a callback function to subscribe to a network event.

Args:
Expand All @@ -118,7 +127,7 @@ def _on_request(self, event_name, callback):
"""
event = NetworkEvent(event_name)

def _callback(event_data):
def _callback(event_data: NetworkEvent) -> None:
request = Request(
network=self,
request_id=event_data.params["request"].get("request", None),
Expand All @@ -132,7 +141,7 @@ def _callback(event_data):
)
callback(request)

callback_id = self.conn.add_callback(event, _callback)
callback_id: int = self.conn.add_callback(event, _callback)

if event_name in self.callbacks:
self.callbacks[event_name].append(callback_id)
Expand All @@ -141,7 +150,13 @@ def _callback(event_data):

return callback_id

def add_request_handler(self, event, callback, url_patterns=None, contexts=None):
def add_request_handler(
self,
event: str,
callback: Callable[["Request"], Any],
Comment thread
cgoldberg marked this conversation as resolved.
Outdated
url_patterns: list[Any] | None = None,
contexts: list[str] | None = None,
) -> int:
"""Add a request handler to the network.

Args:
Expand All @@ -166,15 +181,15 @@ def add_request_handler(self, event, callback, url_patterns=None, contexts=None)
if event_name in self.subscriptions:
self.subscriptions[event_name].append(callback_id)
else:
params = {}
params: dict[str, Any] = {}
params["events"] = [event_name]
self.conn.execute(command_builder("session.subscribe", params))
self.subscriptions[event_name] = [callback_id]

self.callbacks[callback_id] = result["intercept"]
return callback_id

def remove_request_handler(self, event, callback_id):
def remove_request_handler(self, event: str, callback_id: int) -> None:
"""Remove a request handler from the network.

Args:
Expand All @@ -193,25 +208,25 @@ def remove_request_handler(self, event, callback_id):
del self.callbacks[callback_id]
self.subscriptions[event_name].remove(callback_id)
if len(self.subscriptions[event_name]) == 0:
params = {}
params: dict[str, Any] = {}
params["events"] = [event_name]
self.conn.execute(command_builder("session.unsubscribe", params))
del self.subscriptions[event_name]

def clear_request_handlers(self):
def clear_request_handlers(self) -> None:
"""Clear all request handlers from the network."""
for event_name in self.subscriptions:
net_event = NetworkEvent(event_name)
for callback_id in self.subscriptions[event_name]:
self.conn.remove_callback(net_event, callback_id)
self._remove_intercept(self.callbacks[callback_id])
del self.callbacks[callback_id]
params = {}
params: dict[str, Any] = {}
params["events"] = [event_name]
self.conn.execute(command_builder("session.unsubscribe", params))
self.subscriptions = {}

def add_auth_handler(self, username, password):
def add_auth_handler(self, username: str, password: str) -> int:
"""Add an authentication handler to the network.

Args:
Expand All @@ -223,12 +238,12 @@ def add_auth_handler(self, username, password):
"""
event = "auth_required"

def _callback(request):
def _callback(request: Request) -> None:
request._continue_with_auth(username, password)

return self.add_request_handler(event, _callback)

def remove_auth_handler(self, callback_id):
def remove_auth_handler(self, callback_id: int) -> None:
"""Remove an authentication handler from the network.

Args:
Expand All @@ -244,16 +259,16 @@ class Request:
def __init__(
self,
network: Network,
request_id,
body_size=None,
cookies=None,
resource_type=None,
headers=None,
headers_size=None,
method=None,
timings=None,
url=None,
):
request_id: Any,
body_size: int | None = None,
cookies: Any = None,
resource_type: str | None = None,
headers: Any = None,
headers_size: int | None = None,
method: str | None = None,
timings: Any = None,
url: str | None = None,
) -> None:
self.network = network
self.request_id = request_id
self.body_size = body_size
Expand All @@ -265,20 +280,27 @@ def __init__(
self.timings = timings
self.url = url

def fail_request(self):
def fail_request(self) -> None:
"""Fail this request."""
if not self.request_id:
raise ValueError("Request not found.")

params = {"request": self.request_id}
params: dict[str, Any] = {"request": self.request_id}
self.network.conn.execute(command_builder("network.failRequest", params))

def continue_request(self, body=None, method=None, headers=None, cookies=None, url=None):
def continue_request(
self,
body: Any = None,
method: str | None = None,
headers: Any = None,
cookies: Any = None,
url: str | None = None,
) -> None:
"""Continue after intercepting this request."""
if not self.request_id:
raise ValueError("Request not found.")

params = {"request": self.request_id}
params: dict[str, Any] = {"request": self.request_id}
if body is not None:
params["body"] = body
if method is not None:
Expand All @@ -292,7 +314,7 @@ def continue_request(self, body=None, method=None, headers=None, cookies=None, u

self.network.conn.execute(command_builder("network.continueRequest", params))

def _continue_with_auth(self, username=None, password=None):
def _continue_with_auth(self, username: str | None = None, password: str | None = None) -> None:
"""Continue with authentication.

Args:
Expand All @@ -302,7 +324,7 @@ def _continue_with_auth(self, username=None, password=None):
Note:
If username or password is None, it attempts auth with no credentials.
"""
params = {}
params: dict[str, Any] = {}
params["request"] = self.request_id

if not username or not password: # no credentials is valid option
Expand Down
Loading