import unittest from typing import Any, Callable, Awaitable, List, Mapping, Optional from pwo import async_test from kaya import KayaApp, WebSocket, WebSocketMessage def websocket_scope(path: str = '/ws') -> Mapping[str, Any]: return { 'type': 'websocket', 'asgi': {'spec_version': '2.3', 'version': '3.0'}, 'http_version': '1.1', 'scheme': 'ws', 'path': path, 'raw_path': path.encode(), 'query_string': b'', 'root_path': '', 'headers': [], 'client': ('127.0.0.1', 12345), 'server': ('127.0.0.1', 80), 'subprotocols': [], 'extensions': None, } class WebSocketTest(unittest.TestCase): app: KayaApp def setUp(self): self.app = KayaApp() @self.app.websocket('/echo') async def echo(ws: WebSocket) -> None: await ws.accept() async for msg in ws: if msg.kind == 'text': await ws.send_text(f"echo: {msg.data}") elif msg.kind == 'binary': data = msg.data assert isinstance(data, bytes) await ws.send_bytes(data) @self.app.websocket('/room/${room_id}') async def room(ws: WebSocket, room_id: str) -> None: await ws.accept() async for msg in ws: if msg.kind == 'text': await ws.send_text(f"[{room_id}] {msg.data}") @async_test async def test_echo_text(self): sent_messages: List[Mapping[str, Any]] = [] received_messages: List[Mapping[str, Any]] = [] async def receive() -> Mapping[str, Any]: if not received_messages: received_messages.append({'type': 'websocket.connect'}) return received_messages[-1] if len(received_messages) == 1: received_messages.append({'type': 'websocket.receive', 'text': 'hello'}) return received_messages[-1] received_messages.append({'type': 'websocket.disconnect', 'code': 1000}) return received_messages[-1] async def send(message: Mapping[str, Any]) -> None: sent_messages.append(message) await self.app(websocket_scope('/echo'), receive, send) self.assertEqual(sent_messages[0]['type'], 'websocket.accept') self.assertEqual(sent_messages[1]['type'], 'websocket.send') self.assertEqual(sent_messages[1]['text'], 'echo: hello') @async_test async def test_echo_binary(self): sent_messages: List[Mapping[str, Any]] = [] received_messages: List[Mapping[str, Any]] = [] async def receive() -> Mapping[str, Any]: if not received_messages: received_messages.append({'type': 'websocket.connect'}) return received_messages[-1] if len(received_messages) == 1: received_messages.append({'type': 'websocket.receive', 'bytes': b'hello'}) return received_messages[-1] received_messages.append({'type': 'websocket.disconnect', 'code': 1000}) return received_messages[-1] async def send(message: Mapping[str, Any]) -> None: sent_messages.append(message) await self.app(websocket_scope('/echo'), receive, send) self.assertEqual(sent_messages[0]['type'], 'websocket.accept') self.assertEqual(sent_messages[1]['type'], 'websocket.send') self.assertEqual(sent_messages[1]['bytes'], b'hello') @async_test async def test_path_parameter(self): sent_messages: List[Mapping[str, Any]] = [] received_messages: List[Mapping[str, Any]] = [] async def receive() -> Mapping[str, Any]: if not received_messages: received_messages.append({'type': 'websocket.connect'}) return received_messages[-1] if len(received_messages) == 1: received_messages.append({'type': 'websocket.receive', 'text': 'hi'}) return received_messages[-1] received_messages.append({'type': 'websocket.disconnect', 'code': 1000}) return received_messages[-1] async def send(message: Mapping[str, Any]) -> None: sent_messages.append(message) await self.app(websocket_scope('/room/general'), receive, send) self.assertEqual(sent_messages[0]['type'], 'websocket.accept') self.assertEqual(sent_messages[1]['text'], '[general] hi') @async_test async def test_no_handler(self): sent_messages: List[Mapping[str, Any]] = [] async def receive() -> Mapping[str, Any]: return {'type': 'websocket.connect'} async def send(message: Mapping[str, Any]) -> None: sent_messages.append(message) await self.app(websocket_scope('/unknown'), receive, send) self.assertEqual(sent_messages[0]['type'], 'websocket.close') self.assertEqual(sent_messages[0]['code'], 1000) @async_test async def test_close_from_client(self): sent_messages: List[Mapping[str, Any]] = [] received_messages: List[Mapping[str, Any]] = [] async def receive() -> Mapping[str, Any]: if not received_messages: received_messages.append({'type': 'websocket.connect'}) return received_messages[-1] received_messages.append({'type': 'websocket.disconnect', 'code': 1001}) return received_messages[-1] async def send(message: Mapping[str, Any]) -> None: sent_messages.append(message) await self.app(websocket_scope('/echo'), receive, send) self.assertEqual(sent_messages[0]['type'], 'websocket.accept') self.assertEqual(len(sent_messages), 1) class RsgiWebSocketTest(unittest.TestCase): def test_misconfigured_granian(self): from kaya._rsgi import RsgiWebSocket class FakeScope: scheme = 'ws' path = '/ws' query_string = '' headers = {} client = '127.0.0.1:12345' server = '127.0.0.1:80' class FakeProtocol: pass with self.assertRaises(RuntimeError) as ctx: RsgiWebSocket(FakeScope(), FakeProtocol()) # type: ignore[arg-type] self.assertIn('Granian was not configured for websockets', str(ctx.exception))