Add WebSocket support
This commit is contained in:
@@ -3,6 +3,7 @@ from ._http_method import HttpMethod
|
|||||||
from ._http_context import HttpContext
|
from ._http_context import HttpContext
|
||||||
from ._tree import Tree, PathIterator
|
from ._tree import Tree, PathIterator
|
||||||
from ._path_handler import PathHandler
|
from ._path_handler import PathHandler
|
||||||
|
from ._websocket import WebSocket, WebSocketMessage
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
@@ -11,5 +12,7 @@ __all__ = [
|
|||||||
'HttpContext',
|
'HttpContext',
|
||||||
'Tree',
|
'Tree',
|
||||||
'PathHandler',
|
'PathHandler',
|
||||||
'PathIterator'
|
'PathIterator',
|
||||||
|
'WebSocket',
|
||||||
|
'WebSocketMessage'
|
||||||
]
|
]
|
||||||
|
|||||||
+44
-11
@@ -2,36 +2,44 @@ from abc import ABC, abstractmethod
|
|||||||
from asyncio import Queue, AbstractEventLoop
|
from asyncio import Queue, AbstractEventLoop
|
||||||
from asyncio import get_running_loop
|
from asyncio import get_running_loop
|
||||||
from logging import getLogger
|
from logging import getLogger
|
||||||
from typing import Callable, Awaitable, Any, Mapping, Sequence, Optional, Unpack, Tuple, TYPE_CHECKING
|
from typing import Callable, Awaitable, Any, Mapping, Sequence, Optional, Unpack, Tuple, TYPE_CHECKING, cast
|
||||||
from pathlib import Path, PurePath
|
from pathlib import Path, PurePath
|
||||||
from pwo import Maybe, AsyncQueueIterator
|
from pwo import Maybe, AsyncQueueIterator
|
||||||
from hashlib import md5
|
from hashlib import md5
|
||||||
from ._http_context import HttpContext
|
from ._http_context import HttpContext
|
||||||
from ._http_method import HttpMethod
|
from ._http_method import HttpMethod
|
||||||
|
from ._path_handler import Context
|
||||||
from ._types import StrOrStrings
|
from ._types import StrOrStrings
|
||||||
|
from ._websocket import WebSocket
|
||||||
from base64 import b64encode, b64decode
|
from base64 import b64encode, b64decode
|
||||||
from mimetypes import guess_type
|
from mimetypes import guess_type
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from _typeshed import StrOrBytesPath
|
from _typeshed import StrOrBytesPath
|
||||||
try:
|
try:
|
||||||
from ._rsgi import RsgiContext
|
from ._rsgi import RsgiContext, RsgiWebSocket
|
||||||
from granian._granian import RSGIHTTPProtocol, RSGIHTTPScope # type: ignore
|
from granian._granian import ( # type: ignore
|
||||||
|
RSGIHTTPProtocol,
|
||||||
|
RSGIHTTPScope,
|
||||||
|
RSGIWebsocketProtocol,
|
||||||
|
RSGIWebsocketScope as RSGIWebSocketScope,
|
||||||
|
)
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
from ._asgi import AsgiContext
|
from ._asgi import AsgiContext, AsgiWebSocket
|
||||||
from ._tree import Tree
|
from ._tree import Tree
|
||||||
from ._types.asgi import LifespanScope, HTTPScope as ASGIHTTPScope, WebSocketScope
|
from ._types.asgi import LifespanScope, HTTPScope as ASGIHTTPScope, WebSocketScope as ASGIWebSocketScope
|
||||||
|
|
||||||
log = getLogger(__name__)
|
log = getLogger(__name__)
|
||||||
|
|
||||||
type HttpHandler = Callable[[HttpContext, Unpack[Any]], Awaitable[None]]
|
type HttpHandler = Callable[[HttpContext, Unpack[Any]], Awaitable[None]]
|
||||||
|
type WebSocketHandler = Callable[[WebSocket, Unpack[Any]], Awaitable[None]]
|
||||||
|
|
||||||
|
|
||||||
class AbstractBugisApp(ABC):
|
class AbstractBugisApp(ABC):
|
||||||
async def __call__(self,
|
async def __call__(self,
|
||||||
scope: ASGIHTTPScope | WebSocketScope | LifespanScope,
|
scope: ASGIHTTPScope | ASGIWebSocketScope | LifespanScope,
|
||||||
receive: Callable[[], Awaitable[Any]],
|
receive: Callable[[], Awaitable[Any]],
|
||||||
send: Callable[[Mapping[str, Any]], Awaitable[None]]) -> None:
|
send: Callable[[Mapping[str, Any]], Awaitable[None]]) -> None:
|
||||||
loop = get_running_loop()
|
loop = get_running_loop()
|
||||||
@@ -59,6 +67,9 @@ class AbstractBugisApp(ABC):
|
|||||||
elif message['type'] == 'http.disconnect':
|
elif message['type'] == 'http.disconnect':
|
||||||
request_handling.cancel()
|
request_handling.cancel()
|
||||||
break
|
break
|
||||||
|
elif scope['type'] == 'websocket':
|
||||||
|
ws = AsgiWebSocket(scope, receive, send)
|
||||||
|
await self.handle_websocket(ws)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|
||||||
@@ -72,15 +83,23 @@ class AbstractBugisApp(ABC):
|
|||||||
async def handle_request(self, ctx: HttpContext) -> None:
|
async def handle_request(self, ctx: HttpContext) -> None:
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def handle_websocket(self, ws: WebSocket) -> None:
|
||||||
|
raise NotImplementedError()
|
||||||
|
|
||||||
def __rsgi_init__(self, loop: AbstractEventLoop) -> None:
|
def __rsgi_init__(self, loop: AbstractEventLoop) -> None:
|
||||||
self.setup(loop)
|
self.setup(loop)
|
||||||
|
|
||||||
def __rsgi_del__(self, loop: AbstractEventLoop) -> None:
|
def __rsgi_del__(self, loop: AbstractEventLoop) -> None:
|
||||||
self.shutdown(loop)
|
self.shutdown(loop)
|
||||||
|
|
||||||
async def __rsgi__(self, scope: RSGIHTTPScope, protocol: RSGIHTTPProtocol) -> None:
|
async def __rsgi__(self, scope: RSGIHTTPScope | RSGIWebSocketScope, protocol: RSGIHTTPProtocol | RSGIWebsocketProtocol) -> None:
|
||||||
ctx = RsgiContext(scope, protocol)
|
if scope.proto == 'ws':
|
||||||
await self.handle_request(ctx)
|
ws = RsgiWebSocket(scope, protocol) # type: ignore[arg-type]
|
||||||
|
await self.handle_websocket(ws)
|
||||||
|
else:
|
||||||
|
ctx = RsgiContext(scope, protocol) # type: ignore[arg-type]
|
||||||
|
await self.handle_request(ctx)
|
||||||
|
|
||||||
|
|
||||||
class BugisApp(AbstractBugisApp):
|
class BugisApp(AbstractBugisApp):
|
||||||
@@ -98,6 +117,15 @@ class BugisApp(AbstractBugisApp):
|
|||||||
await ctx.send_empty(404)
|
await ctx.send_empty(404)
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
async def handle_websocket(self, ws: WebSocket) -> None:
|
||||||
|
result = self._tree.get_handler(ws.path, HttpMethod.WS)
|
||||||
|
if result is not None:
|
||||||
|
handler, captured = result
|
||||||
|
await handler.handle_request(ws, captured)
|
||||||
|
else:
|
||||||
|
await ws.close(1000)
|
||||||
|
pass
|
||||||
|
|
||||||
def route(self,
|
def route(self,
|
||||||
paths: StrOrStrings,
|
paths: StrOrStrings,
|
||||||
methods: Optional[HttpMethod | Sequence[HttpMethod]] = None,
|
methods: Optional[HttpMethod | Sequence[HttpMethod]] = None,
|
||||||
@@ -120,11 +148,17 @@ class BugisApp(AbstractBugisApp):
|
|||||||
_paths = tuple(paths)
|
_paths = tuple(paths)
|
||||||
for method in _methods:
|
for method in _methods:
|
||||||
for path in _paths:
|
for path in _paths:
|
||||||
self._tree.register(path, method, handler, recursive)
|
self._tree.register(path, method, cast(Callable[[Context, Unpack[Any]], Awaitable[None]], handler), recursive)
|
||||||
return handler
|
return handler
|
||||||
|
|
||||||
return wrapped
|
return wrapped
|
||||||
|
|
||||||
|
def websocket(self, path: str, recursive: bool = False) -> Callable[[WebSocketHandler], WebSocketHandler]:
|
||||||
|
def wrapped(handler: WebSocketHandler) -> WebSocketHandler:
|
||||||
|
self._tree.register(path, HttpMethod.WS, cast(Callable[[Context, Unpack[Any]], Awaitable[None]], handler), recursive)
|
||||||
|
return handler
|
||||||
|
return wrapped
|
||||||
|
|
||||||
def GET(self, path: str, recursive: bool = False) -> Callable[[HttpHandler], HttpHandler]:
|
def GET(self, path: str, recursive: bool = False) -> Callable[[HttpHandler], HttpHandler]:
|
||||||
return self.route(path, (HttpMethod.GET,), recursive)
|
return self.route(path, (HttpMethod.GET,), recursive)
|
||||||
|
|
||||||
@@ -146,4 +180,3 @@ class BugisApp(AbstractBugisApp):
|
|||||||
def PATCH(self, path: str, recursive: bool = False) -> Callable[[HttpHandler], HttpHandler]:
|
def PATCH(self, path: str, recursive: bool = False) -> Callable[[HttpHandler], HttpHandler]:
|
||||||
return self.route(path, (HttpMethod.PATCH,), recursive)
|
return self.route(path, (HttpMethod.PATCH,), recursive)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+62
-1
@@ -17,8 +17,9 @@ from pwo import Maybe
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from ._http_method import HttpMethod
|
from ._http_method import HttpMethod
|
||||||
from ._http_context import HttpContext
|
from ._http_context import HttpContext
|
||||||
|
from ._websocket import WebSocket, WebSocketMessage
|
||||||
from ._types import StrOrStrings
|
from ._types import StrOrStrings
|
||||||
from ._types.asgi import HTTPScope
|
from ._types.asgi import HTTPScope, WebSocketScope
|
||||||
|
|
||||||
|
|
||||||
def decode_headers(headers: Iterable[Tuple[bytes, bytes]]) -> Dict[str, Sequence[str]]:
|
def decode_headers(headers: Iterable[Tuple[bytes, bytes]]) -> Dict[str, Sequence[str]]:
|
||||||
@@ -146,3 +147,63 @@ class AsgiContext(HttpContext):
|
|||||||
'body': '',
|
'body': '',
|
||||||
'more_body': False
|
'more_body': False
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
||||||
|
class AsgiWebSocket(WebSocket):
|
||||||
|
_receive: Callable[[], Awaitable[Any]]
|
||||||
|
_send: Callable[[Mapping[str, Any]], Awaitable[None]]
|
||||||
|
scheme: str
|
||||||
|
path: str
|
||||||
|
query_string: str
|
||||||
|
headers: Mapping[str, Sequence[str]]
|
||||||
|
client: Optional[Tuple[str, int]]
|
||||||
|
server: Optional[Tuple[str, Optional[int]]]
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
scope: WebSocketScope,
|
||||||
|
receive: Callable[[], Awaitable[Any]],
|
||||||
|
send: Callable[[Mapping[str, Any]], Awaitable[None]]):
|
||||||
|
self._receive = receive
|
||||||
|
self._send = send
|
||||||
|
self.path = scope['path']
|
||||||
|
self.query_string = scope['query_string'].decode()
|
||||||
|
self.scheme = scope['scheme']
|
||||||
|
self.client = scope['client']
|
||||||
|
self.server = scope['server']
|
||||||
|
self.headers = decode_headers(scope['headers'])
|
||||||
|
|
||||||
|
async def accept(self) -> None:
|
||||||
|
await self._send({'type': 'websocket.accept'})
|
||||||
|
|
||||||
|
async def receive(self) -> WebSocketMessage:
|
||||||
|
message = await self._receive()
|
||||||
|
message_type: str = message['type']
|
||||||
|
if message_type == 'websocket.connect':
|
||||||
|
return await self.receive()
|
||||||
|
if message_type == 'websocket.receive':
|
||||||
|
if 'text' in message:
|
||||||
|
return WebSocketMessage(kind='text', data=message['text'])
|
||||||
|
elif 'bytes' in message:
|
||||||
|
return WebSocketMessage(kind='binary', data=message['bytes'])
|
||||||
|
else:
|
||||||
|
return WebSocketMessage(kind='close')
|
||||||
|
elif message_type == 'websocket.disconnect':
|
||||||
|
code: int = message.get('code', 1000)
|
||||||
|
return WebSocketMessage(kind='close', data=code)
|
||||||
|
else:
|
||||||
|
return WebSocketMessage(kind='close')
|
||||||
|
|
||||||
|
async def send_text(self, data: str) -> None:
|
||||||
|
await self._send({'type': 'websocket.send', 'text': data})
|
||||||
|
|
||||||
|
async def send_bytes(self, data: bytes) -> None:
|
||||||
|
await self._send({'type': 'websocket.send', 'bytes': data})
|
||||||
|
|
||||||
|
async def close(self, code: int = 1000) -> None:
|
||||||
|
await self._send({'type': 'websocket.close', 'code': code})
|
||||||
|
|
||||||
|
async def __anext__(self) -> WebSocketMessage:
|
||||||
|
message = await self.receive()
|
||||||
|
if message.kind == 'close':
|
||||||
|
raise StopAsyncIteration
|
||||||
|
return message
|
||||||
|
|||||||
@@ -9,3 +9,4 @@ class HttpMethod(StrEnum):
|
|||||||
PUT = 'PUT'
|
PUT = 'PUT'
|
||||||
DELETE = 'DELETE'
|
DELETE = 'DELETE'
|
||||||
PATCH = 'PATCH'
|
PATCH = 'PATCH'
|
||||||
|
WS = 'WS'
|
||||||
|
|||||||
@@ -2,10 +2,12 @@ from abc import ABC, abstractmethod
|
|||||||
from typing import (
|
from typing import (
|
||||||
Sequence,
|
Sequence,
|
||||||
Dict,
|
Dict,
|
||||||
Optional
|
Optional,
|
||||||
|
Union
|
||||||
)
|
)
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from ._http_context import HttpContext
|
from ._http_context import HttpContext
|
||||||
|
from ._websocket import WebSocket
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -18,10 +20,13 @@ class Matches:
|
|||||||
unmatched_paths: Sequence[str] = field(default_factory=list)
|
unmatched_paths: Sequence[str] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
type Context = Union[HttpContext, WebSocket]
|
||||||
|
|
||||||
|
|
||||||
class PathHandler(ABC):
|
class PathHandler(ABC):
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def handle_request(self, ctx: HttpContext, captured: Matches) -> None:
|
async def handle_request(self, ctx: Context, captured: Matches) -> None:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|||||||
+73
-1
@@ -14,12 +14,13 @@ from typing import (
|
|||||||
cast
|
cast
|
||||||
)
|
)
|
||||||
|
|
||||||
from granian._granian import RSGIHTTPProtocol, RSGIHTTPScope # type: ignore[attr-defined]
|
from granian._granian import RSGIHTTPProtocol, RSGIHTTPScope, RSGIWebsocketProtocol, RSGIWebsocketScope, RSGIWebsocketTransport # type: ignore[attr-defined]
|
||||||
from pwo import Maybe
|
from pwo import Maybe
|
||||||
|
|
||||||
from ._types import StrOrStrings
|
from ._types import StrOrStrings
|
||||||
from ._http_context import HttpContext
|
from ._http_context import HttpContext
|
||||||
from ._http_method import HttpMethod
|
from ._http_method import HttpMethod
|
||||||
|
from ._websocket import WebSocket, WebSocketMessage
|
||||||
|
|
||||||
|
|
||||||
class RsgiContext(HttpContext):
|
class RsgiContext(HttpContext):
|
||||||
@@ -105,3 +106,74 @@ class RsgiContext(HttpContext):
|
|||||||
async def send_empty(self, status: int, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
async def send_empty(self, status: int, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
rearranged_headers = Maybe.of_nullable(headers).map(RsgiContext._rearrange_headers).or_else(list())
|
rearranged_headers = Maybe.of_nullable(headers).map(RsgiContext._rearrange_headers).or_else(list())
|
||||||
self.protocol.response_empty(status, rearranged_headers)
|
self.protocol.response_empty(status, rearranged_headers)
|
||||||
|
|
||||||
|
|
||||||
|
class RsgiWebSocket(WebSocket):
|
||||||
|
_protocol: RSGIWebsocketProtocol
|
||||||
|
_transport: Optional[RSGIWebsocketTransport]
|
||||||
|
scheme: str
|
||||||
|
path: str
|
||||||
|
query_string: str
|
||||||
|
headers: Mapping[str, Sequence[str]]
|
||||||
|
client: Optional[Tuple[str, int]]
|
||||||
|
server: Optional[Tuple[str, Optional[int]]]
|
||||||
|
|
||||||
|
def __init__(self, scope: RSGIWebsocketScope, protocol: RSGIWebsocketProtocol):
|
||||||
|
if not hasattr(protocol, 'accept') or not hasattr(protocol, 'close'):
|
||||||
|
raise RuntimeError(
|
||||||
|
'Granian was not configured for websockets; '
|
||||||
|
'ensure Granian is started with websocket support enabled'
|
||||||
|
)
|
||||||
|
self._protocol = protocol
|
||||||
|
self._transport = None
|
||||||
|
self.scheme = scope.scheme
|
||||||
|
self.path = scope.path
|
||||||
|
self.query_string = scope.query_string
|
||||||
|
|
||||||
|
def acc(d: Dict[str, List[str]], t: Tuple[str, str]) -> Dict[str, List[str]]:
|
||||||
|
d.setdefault(t[0].lower(), list()).append(t[1])
|
||||||
|
return d
|
||||||
|
|
||||||
|
fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc)
|
||||||
|
self.headers = reduce(fun, scope.headers.items(), {})
|
||||||
|
self.client = (Maybe.of(scope.client.split(':'))
|
||||||
|
.map(lambda it: (it[0], int(it[1])))
|
||||||
|
.or_else_throw(RuntimeError))
|
||||||
|
self.server = (Maybe.of(scope.server.split(':'))
|
||||||
|
.map(lambda it: (it[0], int(it[1])))
|
||||||
|
.or_else_throw(RuntimeError))
|
||||||
|
|
||||||
|
async def accept(self) -> None:
|
||||||
|
self._transport = await self._protocol.accept()
|
||||||
|
|
||||||
|
async def receive(self) -> WebSocketMessage:
|
||||||
|
if self._transport is None:
|
||||||
|
raise RuntimeError('WebSocket connection has not been accepted yet')
|
||||||
|
message = await self._transport.receive()
|
||||||
|
if message.kind == 0:
|
||||||
|
return WebSocketMessage(kind='close')
|
||||||
|
elif message.kind == 1:
|
||||||
|
return WebSocketMessage(kind='binary', data=message.data)
|
||||||
|
elif message.kind == 2:
|
||||||
|
return WebSocketMessage(kind='text', data=message.data)
|
||||||
|
else:
|
||||||
|
return WebSocketMessage(kind='close')
|
||||||
|
|
||||||
|
async def send_text(self, data: str) -> None:
|
||||||
|
if self._transport is None:
|
||||||
|
raise RuntimeError('WebSocket connection has not been accepted yet')
|
||||||
|
await self._transport.send_str(data)
|
||||||
|
|
||||||
|
async def send_bytes(self, data: bytes) -> None:
|
||||||
|
if self._transport is None:
|
||||||
|
raise RuntimeError('WebSocket connection has not been accepted yet')
|
||||||
|
await self._transport.send_bytes(data)
|
||||||
|
|
||||||
|
async def close(self, code: int = 1000) -> None:
|
||||||
|
self._protocol.close(code)
|
||||||
|
|
||||||
|
async def __anext__(self) -> WebSocketMessage:
|
||||||
|
message = await self.receive()
|
||||||
|
if message.kind == 'close':
|
||||||
|
raise StopAsyncIteration
|
||||||
|
return message
|
||||||
|
|||||||
+3
-3
@@ -18,7 +18,7 @@ from pwo import Maybe, index_of_with_escape
|
|||||||
|
|
||||||
from ._http_context import HttpContext
|
from ._http_context import HttpContext
|
||||||
from ._http_method import HttpMethod
|
from ._http_method import HttpMethod
|
||||||
from ._path_handler import PathHandler
|
from ._path_handler import PathHandler, Context
|
||||||
from ._path_matcher import PathMatcher, IntMatcher, GlobMatcher, StrMatcher, Node
|
from ._path_matcher import PathMatcher, IntMatcher, GlobMatcher, StrMatcher, Node
|
||||||
from ._path_handler import Matches
|
from ._path_handler import Matches
|
||||||
from ._types import NodeType
|
from ._types import NodeType
|
||||||
@@ -135,11 +135,11 @@ class Tree:
|
|||||||
def register(self,
|
def register(self,
|
||||||
path: str,
|
path: str,
|
||||||
method: Optional[HttpMethod],
|
method: Optional[HttpMethod],
|
||||||
callback: Callable[[HttpContext, Unpack[Any]], Awaitable[None]],
|
callback: Callable[[Context, Unpack[Any]], Awaitable[None]],
|
||||||
recursive: bool) -> None:
|
recursive: bool) -> None:
|
||||||
class Handler(PathHandler):
|
class Handler(PathHandler):
|
||||||
|
|
||||||
async def handle_request(self, ctx: HttpContext, captured: Matches) -> None:
|
async def handle_request(self, ctx: Context, captured: Matches) -> None:
|
||||||
args = Maybe.of_nullable(captured.path).map(lambda it: [it]).or_else([])
|
args = Maybe.of_nullable(captured.path).map(lambda it: [it]).or_else([])
|
||||||
await callback(ctx, *args, **captured.kwargs)
|
await callback(ctx, *args, **captured.kwargs)
|
||||||
|
|
||||||
|
|||||||
@@ -24,3 +24,17 @@ class HTTPScope(TypedDict):
|
|||||||
query_string: str
|
query_string: str
|
||||||
headers: Mapping[str, str]
|
headers: Mapping[str, str]
|
||||||
authority: Optional[str]
|
authority: Optional[str]
|
||||||
|
|
||||||
|
|
||||||
|
class WebSocketScope(TypedDict):
|
||||||
|
proto: Literal['ws']
|
||||||
|
rsgi_version: str
|
||||||
|
http_version: str
|
||||||
|
server: str
|
||||||
|
client: str
|
||||||
|
scheme: str
|
||||||
|
method: str
|
||||||
|
path: str
|
||||||
|
query_string: str
|
||||||
|
headers: Mapping[str, str]
|
||||||
|
authority: Optional[str]
|
||||||
@@ -0,0 +1,47 @@
|
|||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import AsyncIterator, Literal, Mapping, Optional, Sequence, Tuple, Union
|
||||||
|
|
||||||
|
type WebSocketData = Optional[Union[str, bytes, int]]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class WebSocketMessage:
|
||||||
|
kind: Literal['text', 'binary', 'close']
|
||||||
|
data: WebSocketData = None
|
||||||
|
|
||||||
|
|
||||||
|
class WebSocket(ABC):
|
||||||
|
path: str
|
||||||
|
query_string: str
|
||||||
|
scheme: str
|
||||||
|
headers: Mapping[str, Sequence[str]]
|
||||||
|
client: Optional[Tuple[str, int]]
|
||||||
|
server: Optional[Tuple[str, Optional[int]]]
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def accept(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def receive(self) -> WebSocketMessage:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def send_text(self, data: str) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def send_bytes(self, data: bytes) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def close(self, code: int = 1000) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def __aiter__(self) -> AsyncIterator[WebSocketMessage]:
|
||||||
|
return self
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def __anext__(self) -> WebSocketMessage:
|
||||||
|
pass
|
||||||
@@ -0,0 +1,175 @@
|
|||||||
|
import unittest
|
||||||
|
from typing import Any, Callable, Awaitable, List, Mapping, Optional
|
||||||
|
from pwo import async_test
|
||||||
|
from kaya import BugisApp, 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: BugisApp
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
self.app = BugisApp()
|
||||||
|
|
||||||
|
@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))
|
||||||
Reference in New Issue
Block a user