382 lines
12 KiB
Python
382 lines
12 KiB
Python
# Copyright (c) Microsoft Corporation.
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
import asyncio
|
|
import re
|
|
from typing import Any, Awaitable, Callable, Literal, Tuple, Union
|
|
|
|
from playwright.async_api import Frame, Page, WebSocketRoute
|
|
from playwright.async_api._generated import Browser
|
|
from tests.server import Server, WebSocketProtocol
|
|
|
|
|
|
async def assert_equal(
|
|
actual_cb: Callable[[], Union[Any, Awaitable[Any]]], expected: Any
|
|
) -> None:
|
|
__tracebackhide__ = True
|
|
start_time = asyncio.get_event_loop().time()
|
|
attempts = 0
|
|
while True:
|
|
actual = actual_cb()
|
|
if asyncio.iscoroutine(actual):
|
|
actual = await actual
|
|
if actual == expected:
|
|
return
|
|
attempts += 1
|
|
if asyncio.get_event_loop().time() - start_time > 5:
|
|
raise TimeoutError(f"Timed out after 10 seconds. Last actual was: {actual}")
|
|
await asyncio.sleep(0.2)
|
|
|
|
|
|
async def setup_ws(
|
|
target: Union[Page, Frame],
|
|
server: Server,
|
|
protocol: Union[Literal["blob"], Literal["arraybuffer"]],
|
|
) -> None:
|
|
await target.goto(server.EMPTY_PAGE)
|
|
await target.evaluate(
|
|
"""({ port, binaryType }) => {
|
|
window.log = [];
|
|
window.ws = new WebSocket('ws://localhost:' + port + '/ws');
|
|
window.ws.binaryType = binaryType;
|
|
window.ws.addEventListener('open', () => window.log.push('open'));
|
|
window.ws.addEventListener('close', event => window.log.push(`close code=${event.code} reason=${event.reason} wasClean=${event.wasClean}`));
|
|
window.ws.addEventListener('error', event => window.log.push(`error`));
|
|
window.ws.addEventListener('message', async event => {
|
|
let data;
|
|
if (typeof event.data === 'string')
|
|
data = event.data;
|
|
else if (event.data instanceof Blob)
|
|
data = 'blob:' + await event.data.text();
|
|
else
|
|
data = 'arraybuffer:' + await (new Blob([event.data])).text();
|
|
window.log.push(`message: data=${data} origin=${event.origin} lastEventId=${event.lastEventId}`);
|
|
});
|
|
window.wsOpened = new Promise(f => window.ws.addEventListener('open', () => f()));
|
|
}""",
|
|
{"port": server.PORT, "binaryType": protocol},
|
|
)
|
|
|
|
|
|
async def test_should_work_with_ws_close(page: Page, server: Server) -> None:
|
|
future: asyncio.Future[WebSocketRoute] = asyncio.Future()
|
|
|
|
def _handle_ws(ws: WebSocketRoute) -> None:
|
|
ws.connect_to_server()
|
|
future.set_result(ws)
|
|
|
|
await page.route_web_socket(re.compile(".*"), _handle_ws)
|
|
|
|
ws_task = server.wait_for_web_socket()
|
|
await setup_ws(page, server, "blob")
|
|
ws = await ws_task
|
|
|
|
route = await future
|
|
route.send("hello")
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=hello origin=ws://localhost:{server.PORT} lastEventId=",
|
|
],
|
|
)
|
|
|
|
closed_promise: asyncio.Future[Tuple[int, str]] = asyncio.Future()
|
|
ws.events.once(
|
|
"close", lambda code, reason: closed_promise.set_result((code, reason))
|
|
)
|
|
await route.close(code=3009, reason="oops")
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=hello origin=ws://localhost:{server.PORT} lastEventId=",
|
|
"close code=3009 reason=oops wasClean=true",
|
|
],
|
|
)
|
|
assert await closed_promise == (3009, "oops")
|
|
|
|
|
|
async def test_should_pattern_match(page: Page, server: Server) -> None:
|
|
await page.route_web_socket(
|
|
re.compile(r".*/ws$"), lambda ws: ws.connect_to_server()
|
|
)
|
|
await page.route_web_socket(
|
|
"**/mock-ws", lambda ws: ws.on_message(lambda message: ws.send("mock-response"))
|
|
)
|
|
|
|
ws_task = server.wait_for_web_socket()
|
|
await page.goto(server.EMPTY_PAGE)
|
|
await page.evaluate(
|
|
"""async ({ port }) => {
|
|
window.log = [];
|
|
window.ws1 = new WebSocket('ws://localhost:' + port + '/ws');
|
|
window.ws1.addEventListener('message', event => window.log.push(`ws1:${event.data}`));
|
|
window.ws2 = new WebSocket('ws://localhost:' + port + '/something/something/mock-ws');
|
|
window.ws2.addEventListener('message', event => window.log.push(`ws2:${event.data}`));
|
|
await Promise.all([
|
|
new Promise(f => window.ws1.addEventListener('open', f)),
|
|
new Promise(f => window.ws2.addEventListener('open', f)),
|
|
]);
|
|
}""",
|
|
{"port": server.PORT},
|
|
)
|
|
|
|
ws = await ws_task
|
|
ws.events.on("message", lambda payload, isBinary: ws.sendMessage(b"response"))
|
|
|
|
await page.evaluate("window.ws1.send('request')")
|
|
await assert_equal(lambda: page.evaluate("window.log"), ["ws1:response"])
|
|
|
|
await page.evaluate("window.ws2.send('request')")
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"), ["ws1:response", "ws2:mock-response"]
|
|
)
|
|
|
|
|
|
async def test_should_work_with_server(page: Page, server: Server) -> None:
|
|
future: asyncio.Future[WebSocketRoute] = asyncio.Future()
|
|
|
|
async def _handle_ws(ws: WebSocketRoute) -> None:
|
|
server = ws.connect_to_server()
|
|
|
|
def _ws_on_message(message: Union[str, bytes]) -> None:
|
|
if message == "to-respond":
|
|
ws.send("response")
|
|
return
|
|
if message == "to-block":
|
|
return
|
|
if message == "to-modify":
|
|
server.send("modified")
|
|
return
|
|
server.send(message)
|
|
|
|
ws.on_message(_ws_on_message)
|
|
|
|
def _server_on_message(message: Union[str, bytes]) -> None:
|
|
if message == "to-block":
|
|
return
|
|
if message == "to-modify":
|
|
ws.send("modified")
|
|
return
|
|
ws.send(message)
|
|
|
|
server.on_message(_server_on_message)
|
|
server.send("fake")
|
|
future.set_result(ws)
|
|
|
|
await page.route_web_socket(re.compile(".*"), _handle_ws)
|
|
ws_task = server.wait_for_web_socket()
|
|
log = []
|
|
|
|
def _once_web_socket_connection(ws: WebSocketProtocol) -> None:
|
|
ws.events.on(
|
|
"message", lambda data, is_binary: log.append(f"message: {data.decode()}")
|
|
)
|
|
ws.events.on(
|
|
"close",
|
|
lambda code, reason: log.append(f"close: code={code} reason={reason}"),
|
|
)
|
|
|
|
server.once_web_socket_connection(_once_web_socket_connection)
|
|
|
|
await setup_ws(page, server, "blob")
|
|
ws = await ws_task
|
|
await assert_equal(lambda: log, ["message: fake"])
|
|
|
|
ws.sendMessage(b"to-modify")
|
|
ws.sendMessage(b"to-block")
|
|
ws.sendMessage(b"pass-server")
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=modified origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=pass-server origin=ws://localhost:{server.PORT} lastEventId=",
|
|
],
|
|
)
|
|
|
|
await page.evaluate(
|
|
"""() => {
|
|
window.ws.send('to-respond');
|
|
window.ws.send('to-modify');
|
|
window.ws.send('to-block');
|
|
window.ws.send('pass-client');
|
|
}"""
|
|
)
|
|
await assert_equal(
|
|
lambda: log, ["message: fake", "message: modified", "message: pass-client"]
|
|
)
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=modified origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=pass-server origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=response origin=ws://localhost:{server.PORT} lastEventId=",
|
|
],
|
|
)
|
|
|
|
route = await future
|
|
route.send("another")
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=modified origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=pass-server origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=response origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=another origin=ws://localhost:{server.PORT} lastEventId=",
|
|
],
|
|
)
|
|
|
|
await page.evaluate(
|
|
"""() => {
|
|
window.ws.send('pass-client-2');
|
|
}"""
|
|
)
|
|
await assert_equal(
|
|
lambda: log,
|
|
[
|
|
"message: fake",
|
|
"message: modified",
|
|
"message: pass-client",
|
|
"message: pass-client-2",
|
|
],
|
|
)
|
|
|
|
await page.evaluate(
|
|
"""() => {
|
|
window.ws.close(3009, 'problem');
|
|
}"""
|
|
)
|
|
await assert_equal(
|
|
lambda: log,
|
|
[
|
|
"message: fake",
|
|
"message: modified",
|
|
"message: pass-client",
|
|
"message: pass-client-2",
|
|
"close: code=3009 reason=problem",
|
|
],
|
|
)
|
|
|
|
|
|
async def test_should_work_without_server(page: Page, server: Server) -> None:
|
|
future: asyncio.Future[WebSocketRoute] = asyncio.Future()
|
|
|
|
async def _handle_ws(ws: WebSocketRoute) -> None:
|
|
def _ws_on_message(message: Union[str, bytes]) -> None:
|
|
if message == "to-respond":
|
|
ws.send("response")
|
|
|
|
ws.on_message(_ws_on_message)
|
|
future.set_result(ws)
|
|
|
|
await page.route_web_socket(re.compile(".*"), _handle_ws)
|
|
await setup_ws(page, server, "blob")
|
|
|
|
await page.evaluate(
|
|
"""async () => {
|
|
await window.wsOpened;
|
|
window.ws.send('to-respond');
|
|
window.ws.send('to-block');
|
|
window.ws.send('to-respond');
|
|
}"""
|
|
)
|
|
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=response origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=response origin=ws://localhost:{server.PORT} lastEventId=",
|
|
],
|
|
)
|
|
|
|
route = await future
|
|
route.send("another")
|
|
# wait for the message to be processed
|
|
await page.wait_for_timeout(100)
|
|
await route.close(code=3008, reason="oops")
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=response origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=response origin=ws://localhost:{server.PORT} lastEventId=",
|
|
f"message: data=another origin=ws://localhost:{server.PORT} lastEventId=",
|
|
"close code=3008 reason=oops wasClean=true",
|
|
],
|
|
)
|
|
|
|
|
|
async def test_should_work_with_base_url(browser: Browser, server: Server) -> None:
|
|
context = await browser.new_context(base_url=f"http://localhost:{server.PORT}")
|
|
page = await context.new_page()
|
|
|
|
async def _handle_ws(ws: WebSocketRoute) -> None:
|
|
ws.on_message(lambda message: ws.send(message))
|
|
|
|
await page.route_web_socket("/ws", _handle_ws)
|
|
await setup_ws(page, server, "blob")
|
|
|
|
await page.evaluate(
|
|
"""async () => {
|
|
await window.wsOpened;
|
|
window.ws.send('echo');
|
|
}"""
|
|
)
|
|
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.log"),
|
|
[
|
|
"open",
|
|
f"message: data=echo origin=ws://localhost:{server.PORT} lastEventId=",
|
|
],
|
|
)
|
|
|
|
|
|
async def test_should_work_with_no_trailing_slash(page: Page, server: Server) -> None:
|
|
log: list[str] = []
|
|
|
|
async def handle_ws(ws: WebSocketRoute) -> None:
|
|
def on_message(message: Union[str, bytes]) -> None:
|
|
assert isinstance(message, str)
|
|
log.append(message)
|
|
ws.send("response")
|
|
|
|
ws.on_message(on_message)
|
|
|
|
# No trailing slash in the route pattern
|
|
await page.route_web_socket(f"ws://localhost:{server.PORT}", handle_ws)
|
|
|
|
await page.goto(server.EMPTY_PAGE)
|
|
await page.evaluate(
|
|
"""({ port }) => {
|
|
window.log = [];
|
|
// No trailing slash in WebSocket URL
|
|
window.ws = new WebSocket('ws://localhost:' + port);
|
|
window.ws.addEventListener('message', event => window.log.push(event.data));
|
|
}""",
|
|
{"port": server.PORT},
|
|
)
|
|
|
|
await assert_equal(
|
|
lambda: page.evaluate("window.ws.readyState"), 1 # WebSocket.OPEN
|
|
)
|
|
await page.evaluate("window.ws.send('query')")
|
|
await assert_equal(lambda: log, ["query"])
|
|
await assert_equal(lambda: page.evaluate("window.log"), ["response"])
|