Загрузить файлы в «venv/Lib/site-packages/uvicorn/protocols/websockets»
This commit is contained in:
BIN
venv/Lib/site-packages/uvicorn/protocols/websockets/__init__.py
Normal file
BIN
venv/Lib/site-packages/uvicorn/protocols/websockets/__init__.py
Normal file
Binary file not shown.
21
venv/Lib/site-packages/uvicorn/protocols/websockets/auto.py
Normal file
21
venv/Lib/site-packages/uvicorn/protocols/websockets/auto.py
Normal file
@@ -0,0 +1,21 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from collections.abc import Callable
|
||||||
|
|
||||||
|
AutoWebSocketsProtocol: Callable[..., asyncio.Protocol] | None
|
||||||
|
try:
|
||||||
|
import websockets # noqa
|
||||||
|
except ImportError: # pragma: no cover
|
||||||
|
try:
|
||||||
|
import wsproto # noqa
|
||||||
|
except ImportError:
|
||||||
|
AutoWebSocketsProtocol = None
|
||||||
|
else:
|
||||||
|
from uvicorn.protocols.websockets.wsproto_impl import WSProtocol
|
||||||
|
|
||||||
|
AutoWebSocketsProtocol = WSProtocol
|
||||||
|
else:
|
||||||
|
from uvicorn.protocols.websockets.websockets_impl import WebSocketProtocol
|
||||||
|
|
||||||
|
AutoWebSocketsProtocol = WebSocketProtocol
|
||||||
@@ -0,0 +1,372 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import http
|
||||||
|
import logging
|
||||||
|
from collections.abc import Sequence
|
||||||
|
from typing import Any, Literal, cast
|
||||||
|
from urllib.parse import unquote
|
||||||
|
|
||||||
|
import websockets
|
||||||
|
import websockets.legacy.handshake
|
||||||
|
from websockets.datastructures import Headers
|
||||||
|
from websockets.exceptions import ConnectionClosed
|
||||||
|
from websockets.extensions.base import ServerExtensionFactory
|
||||||
|
from websockets.extensions.permessage_deflate import ServerPerMessageDeflateFactory
|
||||||
|
from websockets.legacy.server import HTTPResponse
|
||||||
|
from websockets.server import WebSocketServerProtocol
|
||||||
|
from websockets.typing import Subprotocol
|
||||||
|
|
||||||
|
from uvicorn._types import (
|
||||||
|
ASGI3Application,
|
||||||
|
ASGISendEvent,
|
||||||
|
WebSocketConnectEvent,
|
||||||
|
WebSocketDisconnectEvent,
|
||||||
|
WebSocketReceiveEvent,
|
||||||
|
WebSocketScope,
|
||||||
|
)
|
||||||
|
from uvicorn.config import Config
|
||||||
|
from uvicorn.logging import TRACE_LOG_LEVEL
|
||||||
|
from uvicorn.protocols.utils import (
|
||||||
|
ClientDisconnected,
|
||||||
|
get_client_addr,
|
||||||
|
get_local_addr,
|
||||||
|
get_path_with_query_string,
|
||||||
|
get_remote_addr,
|
||||||
|
is_ssl,
|
||||||
|
)
|
||||||
|
from uvicorn.server import ServerState
|
||||||
|
|
||||||
|
|
||||||
|
class Server:
|
||||||
|
closing = False
|
||||||
|
|
||||||
|
def register(self, ws: WebSocketServerProtocol) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def unregister(self, ws: WebSocketServerProtocol) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def is_serving(self) -> bool:
|
||||||
|
return not self.closing
|
||||||
|
|
||||||
|
|
||||||
|
class WebSocketProtocol(WebSocketServerProtocol):
|
||||||
|
extra_headers: list[tuple[str, str]]
|
||||||
|
logger: logging.Logger | logging.LoggerAdapter[Any]
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: Config,
|
||||||
|
server_state: ServerState,
|
||||||
|
app_state: dict[str, Any],
|
||||||
|
_loop: asyncio.AbstractEventLoop | None = None,
|
||||||
|
):
|
||||||
|
if not config.loaded:
|
||||||
|
config.load()
|
||||||
|
|
||||||
|
self.config = config
|
||||||
|
self.app = cast(ASGI3Application, config.loaded_app)
|
||||||
|
self.loop = _loop or asyncio.get_event_loop()
|
||||||
|
self.root_path = config.root_path
|
||||||
|
self.app_state = app_state
|
||||||
|
|
||||||
|
# Shared server state
|
||||||
|
self.connections = server_state.connections
|
||||||
|
self.tasks = server_state.tasks
|
||||||
|
|
||||||
|
# Connection state
|
||||||
|
self.transport: asyncio.Transport = None # type: ignore[assignment]
|
||||||
|
self.server: tuple[str, int | None] | None = None
|
||||||
|
self.client: tuple[str, int] | None = None
|
||||||
|
self.scheme: Literal["wss", "ws"] = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
# Connection events
|
||||||
|
self.scope: WebSocketScope
|
||||||
|
self.handshake_started_event = asyncio.Event()
|
||||||
|
self.handshake_completed_event = asyncio.Event()
|
||||||
|
self.closed_event = asyncio.Event()
|
||||||
|
self.initial_response: HTTPResponse | None = None
|
||||||
|
self.connect_sent = False
|
||||||
|
self.lost_connection_before_handshake = False
|
||||||
|
self.accepted_subprotocol: Subprotocol | None = None
|
||||||
|
|
||||||
|
self.ws_server: Server = Server() # type: ignore[assignment]
|
||||||
|
|
||||||
|
extensions: list[ServerExtensionFactory] = []
|
||||||
|
if self.config.ws_per_message_deflate:
|
||||||
|
extensions.append(ServerPerMessageDeflateFactory())
|
||||||
|
|
||||||
|
super().__init__(
|
||||||
|
ws_handler=self.ws_handler,
|
||||||
|
ws_server=self.ws_server, # type: ignore[arg-type]
|
||||||
|
max_size=self.config.ws_max_size,
|
||||||
|
max_queue=self.config.ws_max_queue,
|
||||||
|
ping_interval=self.config.ws_ping_interval,
|
||||||
|
ping_timeout=self.config.ws_ping_timeout,
|
||||||
|
extensions=extensions,
|
||||||
|
logger=logging.getLogger("uvicorn.error"),
|
||||||
|
)
|
||||||
|
self.server_header = None
|
||||||
|
self.extra_headers = [
|
||||||
|
(name.decode("latin-1"), value.decode("latin-1")) for name, value in server_state.default_headers
|
||||||
|
]
|
||||||
|
|
||||||
|
def connection_made( # type: ignore[override]
|
||||||
|
self, transport: asyncio.Transport
|
||||||
|
) -> None:
|
||||||
|
self.connections.add(self)
|
||||||
|
self.transport = transport
|
||||||
|
self.server = get_local_addr(transport)
|
||||||
|
self.client = get_remote_addr(transport)
|
||||||
|
self.scheme = "wss" if is_ssl(transport) else "ws"
|
||||||
|
|
||||||
|
if self.logger.isEnabledFor(TRACE_LOG_LEVEL):
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket connection made", prefix)
|
||||||
|
|
||||||
|
super().connection_made(transport)
|
||||||
|
|
||||||
|
def connection_lost(self, exc: Exception | None) -> None:
|
||||||
|
self.connections.remove(self)
|
||||||
|
|
||||||
|
if self.logger.isEnabledFor(TRACE_LOG_LEVEL):
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket connection lost", prefix)
|
||||||
|
|
||||||
|
self.lost_connection_before_handshake = not self.handshake_completed_event.is_set()
|
||||||
|
self.handshake_completed_event.set()
|
||||||
|
super().connection_lost(exc)
|
||||||
|
if exc is None:
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def shutdown(self) -> None:
|
||||||
|
self.ws_server.closing = True
|
||||||
|
if self.handshake_completed_event.is_set():
|
||||||
|
self.fail_connection(1012)
|
||||||
|
else:
|
||||||
|
self.send_500_response()
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def on_task_complete(self, task: asyncio.Task[None]) -> None:
|
||||||
|
self.tasks.discard(task)
|
||||||
|
|
||||||
|
async def process_request(self, path: str, request_headers: Headers) -> HTTPResponse | None:
|
||||||
|
"""
|
||||||
|
This hook is called to determine if the websocket should return
|
||||||
|
an HTTP response and close.
|
||||||
|
|
||||||
|
Our behavior here is to start the ASGI application, and then wait
|
||||||
|
for either `accept` or `close` in order to determine if we should
|
||||||
|
close the connection.
|
||||||
|
"""
|
||||||
|
path_portion, _, query_string = path.partition("?")
|
||||||
|
|
||||||
|
websockets.legacy.handshake.check_request(request_headers)
|
||||||
|
|
||||||
|
subprotocols: list[str] = []
|
||||||
|
for header in request_headers.get_all("Sec-WebSocket-Protocol"):
|
||||||
|
subprotocols.extend([token.strip() for token in header.split(",")])
|
||||||
|
|
||||||
|
asgi_headers = [
|
||||||
|
(name.encode("ascii"), value.encode("ascii", errors="surrogateescape"))
|
||||||
|
for name, value in request_headers.raw_items()
|
||||||
|
]
|
||||||
|
path = unquote(path_portion)
|
||||||
|
full_path = self.root_path + path
|
||||||
|
full_raw_path = self.root_path.encode("ascii") + path_portion.encode("ascii")
|
||||||
|
|
||||||
|
self.scope = {
|
||||||
|
"type": "websocket",
|
||||||
|
"asgi": {"version": self.config.asgi_version, "spec_version": "2.4"},
|
||||||
|
"http_version": "1.1",
|
||||||
|
"scheme": self.scheme,
|
||||||
|
"server": self.server,
|
||||||
|
"client": self.client,
|
||||||
|
"root_path": self.root_path,
|
||||||
|
"path": full_path,
|
||||||
|
"raw_path": full_raw_path,
|
||||||
|
"query_string": query_string.encode("ascii"),
|
||||||
|
"headers": asgi_headers,
|
||||||
|
"subprotocols": subprotocols,
|
||||||
|
"state": self.app_state.copy(),
|
||||||
|
"extensions": {"websocket.http.response": {}},
|
||||||
|
}
|
||||||
|
task = self.loop.create_task(self.run_asgi())
|
||||||
|
task.add_done_callback(self.on_task_complete)
|
||||||
|
self.tasks.add(task)
|
||||||
|
await self.handshake_started_event.wait()
|
||||||
|
return self.initial_response
|
||||||
|
|
||||||
|
def process_subprotocol(
|
||||||
|
self, headers: Headers, available_subprotocols: Sequence[Subprotocol] | None
|
||||||
|
) -> Subprotocol | None:
|
||||||
|
"""
|
||||||
|
We override the standard 'process_subprotocol' behavior here so that
|
||||||
|
we return whatever subprotocol is sent in the 'accept' message.
|
||||||
|
"""
|
||||||
|
return self.accepted_subprotocol
|
||||||
|
|
||||||
|
def send_500_response(self) -> None:
|
||||||
|
msg = b"Internal Server Error"
|
||||||
|
content = [
|
||||||
|
b"HTTP/1.1 500 Internal Server Error\r\ncontent-type: text/plain; charset=utf-8\r\n",
|
||||||
|
b"content-length: " + str(len(msg)).encode("ascii") + b"\r\n",
|
||||||
|
b"connection: close\r\n",
|
||||||
|
b"\r\n",
|
||||||
|
msg,
|
||||||
|
]
|
||||||
|
self.transport.write(b"".join(content))
|
||||||
|
# Allow handler task to terminate cleanly, as websockets doesn't cancel it by
|
||||||
|
# itself (see https://github.com/Kludex/uvicorn/issues/920)
|
||||||
|
self.handshake_started_event.set()
|
||||||
|
|
||||||
|
async def ws_handler(self, protocol: WebSocketServerProtocol, path: str) -> Any: # type: ignore[override]
|
||||||
|
"""
|
||||||
|
This is the main handler function for the 'websockets' implementation
|
||||||
|
to call into. We just wait for close then return, and instead allow
|
||||||
|
'send' and 'receive' events to drive the flow.
|
||||||
|
"""
|
||||||
|
self.handshake_completed_event.set()
|
||||||
|
await self.wait_closed()
|
||||||
|
|
||||||
|
async def run_asgi(self) -> None:
|
||||||
|
"""
|
||||||
|
Wrapper around the ASGI callable, handling exceptions and unexpected
|
||||||
|
termination states.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
result = await self.app(self.scope, self.asgi_receive, self.asgi_send) # type: ignore[func-returns-value]
|
||||||
|
except ClientDisconnected: # pragma: full coverage
|
||||||
|
self.closed_event.set()
|
||||||
|
except BaseException:
|
||||||
|
self.closed_event.set()
|
||||||
|
self.logger.exception("Exception in ASGI application\n")
|
||||||
|
if not self.handshake_started_event.is_set():
|
||||||
|
self.send_500_response()
|
||||||
|
else:
|
||||||
|
await self.handshake_completed_event.wait()
|
||||||
|
else:
|
||||||
|
self.closed_event.set()
|
||||||
|
if not self.handshake_started_event.is_set():
|
||||||
|
self.logger.error("ASGI callable returned without sending handshake.")
|
||||||
|
self.send_500_response()
|
||||||
|
elif result is not None:
|
||||||
|
self.logger.error("ASGI callable should return None, but returned '%s'.", result)
|
||||||
|
await self.handshake_completed_event.wait()
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
async def asgi_send(self, message: ASGISendEvent) -> None:
|
||||||
|
if not self.handshake_started_event.is_set():
|
||||||
|
if message["type"] == "websocket.accept":
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" [accepted]',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
)
|
||||||
|
self.initial_response = None
|
||||||
|
self.accepted_subprotocol = cast(Subprotocol | None, message.get("subprotocol"))
|
||||||
|
if "headers" in message:
|
||||||
|
self.extra_headers.extend(
|
||||||
|
# ASGI spec requires bytes
|
||||||
|
# But for compatibility we need to convert it to strings
|
||||||
|
(name.decode("latin-1"), value.decode("latin-1"))
|
||||||
|
for name, value in message["headers"]
|
||||||
|
)
|
||||||
|
self.handshake_started_event.set()
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.close":
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" 403',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
)
|
||||||
|
self.initial_response = (http.HTTPStatus.FORBIDDEN, [], b"")
|
||||||
|
self.handshake_started_event.set()
|
||||||
|
self.closed_event.set()
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.http.response.start":
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" %d',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
message["status"],
|
||||||
|
)
|
||||||
|
# websockets requires the status to be an enum. look it up.
|
||||||
|
status = http.HTTPStatus(message["status"])
|
||||||
|
headers = [
|
||||||
|
(name.decode("latin-1"), value.decode("latin-1")) for name, value in message.get("headers", [])
|
||||||
|
]
|
||||||
|
self.initial_response = (status, headers, b"")
|
||||||
|
self.handshake_started_event.set()
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Expected ASGI message 'websocket.accept', 'websocket.close', "
|
||||||
|
f"or 'websocket.http.response.start' but got '{message['type']}'."
|
||||||
|
)
|
||||||
|
|
||||||
|
elif not self.closed_event.is_set() and self.initial_response is None:
|
||||||
|
await self.handshake_completed_event.wait()
|
||||||
|
|
||||||
|
try:
|
||||||
|
if message["type"] == "websocket.send":
|
||||||
|
bytes_data = message.get("bytes")
|
||||||
|
text_data = message.get("text")
|
||||||
|
data = text_data if bytes_data is None else bytes_data
|
||||||
|
await self.send(data) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.close":
|
||||||
|
code = message.get("code", 1000)
|
||||||
|
reason = message.get("reason", "") or ""
|
||||||
|
await self.close(code, reason)
|
||||||
|
self.closed_event.set()
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Expected ASGI message 'websocket.send' or 'websocket.close', but got '{message['type']}'."
|
||||||
|
)
|
||||||
|
except ConnectionClosed as exc:
|
||||||
|
raise ClientDisconnected from exc
|
||||||
|
|
||||||
|
elif self.initial_response is not None:
|
||||||
|
if message["type"] == "websocket.http.response.body":
|
||||||
|
body = self.initial_response[2] + message["body"]
|
||||||
|
self.initial_response = self.initial_response[:2] + (body,)
|
||||||
|
if not message.get("more_body", False):
|
||||||
|
self.closed_event.set()
|
||||||
|
else:
|
||||||
|
raise RuntimeError(f"Expected ASGI message 'websocket.http.response.body' but got '{message['type']}'.")
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Unexpected ASGI message '{message['type']}', after sending 'websocket.close' "
|
||||||
|
"or response already completed."
|
||||||
|
)
|
||||||
|
|
||||||
|
async def asgi_receive(self) -> WebSocketDisconnectEvent | WebSocketConnectEvent | WebSocketReceiveEvent:
|
||||||
|
if not self.connect_sent:
|
||||||
|
self.connect_sent = True
|
||||||
|
return {"type": "websocket.connect"}
|
||||||
|
|
||||||
|
await self.handshake_completed_event.wait()
|
||||||
|
|
||||||
|
if self.lost_connection_before_handshake:
|
||||||
|
# If the handshake failed or the app closed before handshake completion,
|
||||||
|
# use 1006 Abnormal Closure.
|
||||||
|
return {"type": "websocket.disconnect", "code": 1006}
|
||||||
|
|
||||||
|
if self.closed_event.is_set():
|
||||||
|
return {"type": "websocket.disconnect", "code": 1005}
|
||||||
|
|
||||||
|
try:
|
||||||
|
data = await self.recv()
|
||||||
|
except ConnectionClosed:
|
||||||
|
self.closed_event.set()
|
||||||
|
if self.ws_server.closing:
|
||||||
|
return {"type": "websocket.disconnect", "code": 1012}
|
||||||
|
return {"type": "websocket.disconnect", "code": self.close_code or 1005, "reason": self.close_reason}
|
||||||
|
|
||||||
|
if isinstance(data, str):
|
||||||
|
return {"type": "websocket.receive", "text": data}
|
||||||
|
return {"type": "websocket.receive", "bytes": data}
|
||||||
@@ -0,0 +1,477 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import random
|
||||||
|
import struct
|
||||||
|
import sys
|
||||||
|
from asyncio import TimerHandle
|
||||||
|
from asyncio.transports import BaseTransport, Transport
|
||||||
|
from http import HTTPStatus
|
||||||
|
from typing import Any, Literal, cast
|
||||||
|
from urllib.parse import unquote
|
||||||
|
|
||||||
|
from websockets.exceptions import InvalidState
|
||||||
|
from websockets.extensions.permessage_deflate import ServerPerMessageDeflateFactory
|
||||||
|
from websockets.frames import Frame, Opcode
|
||||||
|
from websockets.http11 import Request
|
||||||
|
from websockets.server import ServerProtocol
|
||||||
|
|
||||||
|
from uvicorn._types import (
|
||||||
|
ASGIReceiveEvent,
|
||||||
|
ASGISendEvent,
|
||||||
|
WebSocketScope,
|
||||||
|
)
|
||||||
|
from uvicorn.config import Config
|
||||||
|
from uvicorn.logging import TRACE_LOG_LEVEL
|
||||||
|
from uvicorn.protocols.utils import (
|
||||||
|
ClientDisconnected,
|
||||||
|
get_client_addr,
|
||||||
|
get_local_addr,
|
||||||
|
get_path_with_query_string,
|
||||||
|
get_remote_addr,
|
||||||
|
is_ssl,
|
||||||
|
)
|
||||||
|
from uvicorn.server import ServerState
|
||||||
|
|
||||||
|
if sys.version_info >= (3, 11): # pragma: no cover
|
||||||
|
from typing import assert_never
|
||||||
|
else: # pragma: no cover
|
||||||
|
from typing_extensions import assert_never
|
||||||
|
|
||||||
|
|
||||||
|
class WebSocketsSansIOProtocol(asyncio.Protocol):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: Config,
|
||||||
|
server_state: ServerState,
|
||||||
|
app_state: dict[str, Any],
|
||||||
|
_loop: asyncio.AbstractEventLoop | None = None,
|
||||||
|
) -> None:
|
||||||
|
if not config.loaded:
|
||||||
|
config.load() # pragma: no cover
|
||||||
|
|
||||||
|
self.config = config
|
||||||
|
self.app = config.loaded_app
|
||||||
|
self.loop = _loop or asyncio.get_event_loop()
|
||||||
|
self.logger = logging.getLogger("uvicorn.error")
|
||||||
|
self.root_path = config.root_path
|
||||||
|
self.app_state = app_state
|
||||||
|
|
||||||
|
# Shared server state
|
||||||
|
self.connections = server_state.connections
|
||||||
|
self.tasks = server_state.tasks
|
||||||
|
self.default_headers = server_state.default_headers
|
||||||
|
|
||||||
|
# Connection state
|
||||||
|
self.transport: asyncio.Transport = None # type: ignore[assignment]
|
||||||
|
self.server: tuple[str, int | None] | None = None
|
||||||
|
self.client: tuple[str, int] | None = None
|
||||||
|
self.scheme: Literal["wss", "ws"] = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
# WebSocket state
|
||||||
|
self.queue: asyncio.Queue[ASGIReceiveEvent] = asyncio.Queue()
|
||||||
|
self.handshake_initiated = False
|
||||||
|
self.handshake_complete = False
|
||||||
|
self.close_sent = False
|
||||||
|
self.initial_response: tuple[int, list[tuple[str, str]], bytes] | None = None
|
||||||
|
|
||||||
|
extensions = []
|
||||||
|
if self.config.ws_per_message_deflate:
|
||||||
|
extensions = [
|
||||||
|
ServerPerMessageDeflateFactory(
|
||||||
|
server_max_window_bits=12,
|
||||||
|
client_max_window_bits=12,
|
||||||
|
compress_settings={"memLevel": 5},
|
||||||
|
)
|
||||||
|
]
|
||||||
|
self.conn = ServerProtocol(
|
||||||
|
extensions=extensions,
|
||||||
|
max_size=self.config.ws_max_size,
|
||||||
|
logger=logging.getLogger("uvicorn.error"),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.read_paused = False
|
||||||
|
self.writable = asyncio.Event()
|
||||||
|
self.writable.set()
|
||||||
|
|
||||||
|
# Keepalive state
|
||||||
|
self.ping_interval = config.ws_ping_interval
|
||||||
|
self.ping_timeout = config.ws_ping_timeout
|
||||||
|
self.ping_timer: TimerHandle | None = None
|
||||||
|
self.pong_timer: TimerHandle | None = None
|
||||||
|
self.pending_ping_payload: bytes | None = None
|
||||||
|
self.ping_sent_at: float = 0.0
|
||||||
|
self.last_ping_rtt: float = 0.0
|
||||||
|
|
||||||
|
# Buffers
|
||||||
|
self.bytes = bytearray()
|
||||||
|
|
||||||
|
def connection_made(self, transport: BaseTransport) -> None:
|
||||||
|
"""Called when a connection is made."""
|
||||||
|
transport = cast(Transport, transport)
|
||||||
|
self.connections.add(self)
|
||||||
|
self.transport = transport
|
||||||
|
self.server = get_local_addr(transport)
|
||||||
|
self.client = get_remote_addr(transport)
|
||||||
|
self.scheme = "wss" if is_ssl(transport) else "ws"
|
||||||
|
|
||||||
|
if self.logger.level <= TRACE_LOG_LEVEL:
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket connection made", prefix)
|
||||||
|
|
||||||
|
def connection_lost(self, exc: Exception | None) -> None:
|
||||||
|
self.stop_keepalive()
|
||||||
|
code = 1005 if self.handshake_complete else 1006
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": code})
|
||||||
|
self.connections.remove(self)
|
||||||
|
|
||||||
|
if self.logger.level <= TRACE_LOG_LEVEL:
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket connection lost", prefix)
|
||||||
|
|
||||||
|
self.handshake_complete = True
|
||||||
|
if exc is None:
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def eof_received(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def shutdown(self) -> None:
|
||||||
|
self.stop_keepalive()
|
||||||
|
if self.handshake_complete:
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": 1012})
|
||||||
|
self.conn.send_close(1012)
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
else:
|
||||||
|
self.send_500_response()
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def data_received(self, data: bytes) -> None:
|
||||||
|
self.conn.receive_data(data)
|
||||||
|
if self.conn.parser_exc is not None: # pragma: no cover
|
||||||
|
self.handle_parser_exception()
|
||||||
|
return
|
||||||
|
self.handle_events()
|
||||||
|
|
||||||
|
def handle_events(self) -> None:
|
||||||
|
for event in self.conn.events_received():
|
||||||
|
if isinstance(event, Request):
|
||||||
|
self.handle_connect(event)
|
||||||
|
if isinstance(event, Frame):
|
||||||
|
if event.opcode == Opcode.CONT:
|
||||||
|
self.handle_cont(event) # pragma: no cover
|
||||||
|
elif event.opcode == Opcode.TEXT:
|
||||||
|
self.handle_text(event)
|
||||||
|
elif event.opcode == Opcode.BINARY:
|
||||||
|
self.handle_bytes(event)
|
||||||
|
elif event.opcode == Opcode.PING:
|
||||||
|
self.handle_ping()
|
||||||
|
elif event.opcode == Opcode.PONG:
|
||||||
|
self.handle_pong(event)
|
||||||
|
elif event.opcode == Opcode.CLOSE:
|
||||||
|
self.handle_close(event)
|
||||||
|
else:
|
||||||
|
assert_never(event.opcode) # pragma: no cover
|
||||||
|
|
||||||
|
# Event handlers
|
||||||
|
|
||||||
|
def handle_connect(self, event: Request) -> None:
|
||||||
|
self.request = event
|
||||||
|
self.response = self.conn.accept(event)
|
||||||
|
self.handshake_initiated = True
|
||||||
|
if self.response.status_code != 101:
|
||||||
|
self.handshake_complete = True
|
||||||
|
self.close_sent = True
|
||||||
|
self.conn.send_response(self.response)
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
self.transport.close()
|
||||||
|
return
|
||||||
|
|
||||||
|
headers = [
|
||||||
|
(key.encode("ascii"), value.encode("ascii", errors="surrogateescape"))
|
||||||
|
for key, value in event.headers.raw_items()
|
||||||
|
]
|
||||||
|
raw_path, _, query_string = event.path.partition("?")
|
||||||
|
self.scope: WebSocketScope = {
|
||||||
|
"type": "websocket",
|
||||||
|
"asgi": {"version": self.config.asgi_version, "spec_version": "2.4"},
|
||||||
|
"http_version": "1.1",
|
||||||
|
"scheme": self.scheme,
|
||||||
|
"server": self.server,
|
||||||
|
"client": self.client,
|
||||||
|
"root_path": self.root_path,
|
||||||
|
"path": self.root_path + unquote(raw_path),
|
||||||
|
"raw_path": self.root_path.encode("ascii") + raw_path.encode("ascii"),
|
||||||
|
"query_string": query_string.encode("ascii"),
|
||||||
|
"headers": headers,
|
||||||
|
"subprotocols": event.headers.get_all("Sec-WebSocket-Protocol"),
|
||||||
|
"state": self.app_state.copy(),
|
||||||
|
"extensions": {"websocket.http.response": {}},
|
||||||
|
}
|
||||||
|
self.queue.put_nowait({"type": "websocket.connect"})
|
||||||
|
task = self.loop.create_task(self.run_asgi())
|
||||||
|
task.add_done_callback(self.on_task_complete)
|
||||||
|
self.tasks.add(task)
|
||||||
|
|
||||||
|
def handle_cont(self, event: Frame) -> None:
|
||||||
|
self.bytes.extend(event.data)
|
||||||
|
if event.fin:
|
||||||
|
self.send_receive_event_to_app()
|
||||||
|
|
||||||
|
def handle_text(self, event: Frame) -> None:
|
||||||
|
self.bytes = bytearray(event.data)
|
||||||
|
self.curr_msg_data_type: Literal["text", "bytes"] = "text"
|
||||||
|
if event.fin:
|
||||||
|
self.send_receive_event_to_app()
|
||||||
|
|
||||||
|
def handle_bytes(self, event: Frame) -> None:
|
||||||
|
self.bytes = bytearray(event.data)
|
||||||
|
self.curr_msg_data_type = "bytes"
|
||||||
|
if event.fin:
|
||||||
|
self.send_receive_event_to_app()
|
||||||
|
|
||||||
|
def send_receive_event_to_app(self) -> None:
|
||||||
|
if self.curr_msg_data_type == "text":
|
||||||
|
try:
|
||||||
|
self.queue.put_nowait({"type": "websocket.receive", "text": self.bytes.decode()})
|
||||||
|
except UnicodeDecodeError: # pragma: no cover
|
||||||
|
self.logger.exception("Invalid UTF-8 sequence received from client.")
|
||||||
|
self.conn.send_close(1007)
|
||||||
|
self.handle_parser_exception()
|
||||||
|
return
|
||||||
|
else:
|
||||||
|
self.queue.put_nowait({"type": "websocket.receive", "bytes": bytes(self.bytes)})
|
||||||
|
if not self.read_paused:
|
||||||
|
self.read_paused = True
|
||||||
|
self.transport.pause_reading()
|
||||||
|
|
||||||
|
def handle_ping(self) -> None:
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
|
||||||
|
def handle_pong(self, event: Frame) -> None:
|
||||||
|
# Ignore unsolicited pongs and stale pongs whose payload doesn't match the ping currently in flight
|
||||||
|
if self.pending_ping_payload is None or bytes(event.data) != self.pending_ping_payload:
|
||||||
|
return # pragma: no cover
|
||||||
|
|
||||||
|
self.last_ping_rtt = self.loop.time() - self.ping_sent_at
|
||||||
|
self.pending_ping_payload = None
|
||||||
|
# The peer answered in time; cancel the pong deadline and chain the next ping. This `schedule_ping()` call is
|
||||||
|
# what keeps the keepalive loop running when ping_timeout is set. When ping_timeout is None the next ping is
|
||||||
|
# already scheduled by `send_keepalive_ping`, so we must not schedule a duplicate here.
|
||||||
|
if self.pong_timer is not None:
|
||||||
|
self.pong_timer.cancel()
|
||||||
|
self.pong_timer = None
|
||||||
|
self.schedule_ping()
|
||||||
|
|
||||||
|
def start_keepalive(self) -> None:
|
||||||
|
if self.ping_interval is not None and self.ping_interval > 0:
|
||||||
|
self.schedule_ping()
|
||||||
|
|
||||||
|
def stop_keepalive(self) -> None:
|
||||||
|
if self.ping_timer is not None:
|
||||||
|
self.ping_timer.cancel()
|
||||||
|
self.ping_timer = None
|
||||||
|
if self.pong_timer is not None: # pragma: no cover
|
||||||
|
self.pong_timer.cancel()
|
||||||
|
self.pong_timer = None
|
||||||
|
self.pending_ping_payload = None
|
||||||
|
|
||||||
|
def schedule_ping(self) -> None:
|
||||||
|
assert self.ping_interval is not None
|
||||||
|
delay = max(0.0, self.ping_interval - self.last_ping_rtt)
|
||||||
|
self.ping_timer = self.loop.call_later(delay, self.send_keepalive_ping)
|
||||||
|
|
||||||
|
def send_keepalive_ping(self) -> None:
|
||||||
|
self.ping_timer = None
|
||||||
|
if self.close_sent or self.transport.is_closing(): # pragma: no cover
|
||||||
|
return
|
||||||
|
# Random 4-byte payload identifies this ping; `handle_pong` uses it to ignore stale or unsolicited pongs.
|
||||||
|
# See https://github.com/python-websockets/websockets/blob/4d229bf9f583d593aa103287aee0a77c9fbc3a79/src/websockets/asyncio/connection.py#L624
|
||||||
|
self.pending_ping_payload = struct.pack("!I", random.getrandbits(32))
|
||||||
|
self.ping_sent_at = self.loop.time()
|
||||||
|
self.conn.send_ping(self.pending_ping_payload)
|
||||||
|
self.transport.write(b"".join(self.conn.data_to_send()))
|
||||||
|
if self.ping_timeout is not None:
|
||||||
|
self.pong_timer = self.loop.call_later(self.ping_timeout, self.keepalive_timeout)
|
||||||
|
else: # pragma: no cover
|
||||||
|
self.schedule_ping()
|
||||||
|
|
||||||
|
def keepalive_timeout(self) -> None:
|
||||||
|
self.pong_timer = None
|
||||||
|
self.pending_ping_payload = None
|
||||||
|
if self.close_sent or self.transport.is_closing(): # pragma: no cover
|
||||||
|
return
|
||||||
|
if self.logger.level <= TRACE_LOG_LEVEL:
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket keepalive ping timeout", prefix)
|
||||||
|
self.conn.fail(1011, "keepalive ping timeout")
|
||||||
|
self.transport.write(b"".join(self.conn.data_to_send()))
|
||||||
|
self.close_sent = True
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def handle_close(self, event: Frame) -> None:
|
||||||
|
if not self.close_sent and not self.transport.is_closing():
|
||||||
|
assert self.conn.close_rcvd is not None
|
||||||
|
code = self.conn.close_rcvd.code
|
||||||
|
reason = self.conn.close_rcvd.reason
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": code, "reason": reason})
|
||||||
|
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def handle_parser_exception(self) -> None: # pragma: no cover
|
||||||
|
assert self.conn.close_sent is not None
|
||||||
|
code = self.conn.close_sent.code
|
||||||
|
reason = self.conn.close_sent.reason
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": code, "reason": reason})
|
||||||
|
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
self.close_sent = True
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def on_task_complete(self, task: asyncio.Task[None]) -> None:
|
||||||
|
self.tasks.discard(task)
|
||||||
|
|
||||||
|
async def run_asgi(self) -> None:
|
||||||
|
try:
|
||||||
|
result = await self.app(self.scope, self.receive, self.send)
|
||||||
|
except ClientDisconnected:
|
||||||
|
pass # pragma: full coverage
|
||||||
|
except BaseException:
|
||||||
|
self.logger.exception("Exception in ASGI application\n")
|
||||||
|
self.send_500_response()
|
||||||
|
else:
|
||||||
|
if not self.handshake_complete:
|
||||||
|
self.logger.error("ASGI callable returned without completing handshake.")
|
||||||
|
self.send_500_response()
|
||||||
|
elif result is not None:
|
||||||
|
self.logger.error("ASGI callable should return None, but returned '%s'.", result)
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def send_500_response(self) -> None:
|
||||||
|
if self.initial_response or self.handshake_complete:
|
||||||
|
return
|
||||||
|
response = self.conn.reject(500, "Internal Server Error")
|
||||||
|
self.conn.send_response(response)
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
|
||||||
|
async def send(self, message: ASGISendEvent) -> None:
|
||||||
|
await self.writable.wait()
|
||||||
|
|
||||||
|
if not self.handshake_complete and self.initial_response is None:
|
||||||
|
if message["type"] == "websocket.accept":
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" [accepted]',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
)
|
||||||
|
headers = [
|
||||||
|
(name.decode("latin-1").lower(), value.decode("latin-1"))
|
||||||
|
for name, value in (self.default_headers + list(message.get("headers", [])))
|
||||||
|
]
|
||||||
|
accepted_subprotocol = message.get("subprotocol")
|
||||||
|
if accepted_subprotocol:
|
||||||
|
headers.append(("Sec-WebSocket-Protocol", accepted_subprotocol))
|
||||||
|
self.response.headers.update(headers)
|
||||||
|
|
||||||
|
if not self.transport.is_closing():
|
||||||
|
self.handshake_complete = True
|
||||||
|
self.conn.send_response(self.response)
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
self.start_keepalive()
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.close":
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": 1006})
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" 403',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
)
|
||||||
|
response = self.conn.reject(HTTPStatus.FORBIDDEN, "")
|
||||||
|
self.conn.send_response(response)
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.close_sent = True
|
||||||
|
self.handshake_complete = True
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
self.transport.close()
|
||||||
|
elif message["type"] == "websocket.http.response.start" and self.initial_response is None:
|
||||||
|
if not (100 <= message["status"] < 600):
|
||||||
|
raise RuntimeError("Invalid HTTP status code '%d' in response." % message["status"])
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" %d',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
message["status"],
|
||||||
|
)
|
||||||
|
headers = [
|
||||||
|
(name.decode("latin-1"), value.decode("latin-1"))
|
||||||
|
for name, value in list(message.get("headers", []))
|
||||||
|
]
|
||||||
|
self.initial_response = (message["status"], headers, b"")
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Expected ASGI message 'websocket.accept', 'websocket.close' "
|
||||||
|
f"or 'websocket.http.response.start' but got '{message['type']}'."
|
||||||
|
)
|
||||||
|
|
||||||
|
elif not self.close_sent and self.initial_response is None:
|
||||||
|
try:
|
||||||
|
if message["type"] == "websocket.send":
|
||||||
|
bytes_data = message.get("bytes")
|
||||||
|
text_data = message.get("text")
|
||||||
|
if bytes_data is not None:
|
||||||
|
self.conn.send_binary(bytes_data)
|
||||||
|
elif text_data is not None:
|
||||||
|
self.conn.send_text(text_data.encode())
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.close":
|
||||||
|
if not self.transport.is_closing():
|
||||||
|
code = message.get("code", 1000)
|
||||||
|
reason = message.get("reason", "") or ""
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": code, "reason": reason})
|
||||||
|
self.conn.send_close(code, reason)
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
self.close_sent = True
|
||||||
|
self.transport.close()
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Expected ASGI message 'websocket.send' or 'websocket.close', but got '{message['type']}'."
|
||||||
|
)
|
||||||
|
except InvalidState:
|
||||||
|
raise ClientDisconnected()
|
||||||
|
elif self.initial_response is not None:
|
||||||
|
if message["type"] == "websocket.http.response.body":
|
||||||
|
body = self.initial_response[2] + message["body"]
|
||||||
|
self.initial_response = self.initial_response[:2] + (body,)
|
||||||
|
if not message.get("more_body", False):
|
||||||
|
response = self.conn.reject(self.initial_response[0], body.decode())
|
||||||
|
response.headers.update(self.initial_response[1])
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": 1006})
|
||||||
|
self.conn.send_response(response)
|
||||||
|
output = self.conn.data_to_send()
|
||||||
|
self.close_sent = True
|
||||||
|
self.transport.write(b"".join(output))
|
||||||
|
self.transport.close()
|
||||||
|
else: # pragma: no cover
|
||||||
|
raise RuntimeError(f"Expected ASGI message 'websocket.http.response.body' but got '{message['type']}'.")
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(f"Unexpected ASGI message '{message['type']}', after sending 'websocket.close'.")
|
||||||
|
|
||||||
|
async def receive(self) -> ASGIReceiveEvent:
|
||||||
|
message = await self.queue.get()
|
||||||
|
if self.read_paused and self.queue.empty():
|
||||||
|
self.read_paused = False
|
||||||
|
self.transport.resume_reading()
|
||||||
|
return message
|
||||||
@@ -0,0 +1,458 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import random
|
||||||
|
import struct
|
||||||
|
from asyncio import TimerHandle
|
||||||
|
from io import BytesIO, StringIO
|
||||||
|
from typing import Any, Literal, cast
|
||||||
|
from urllib.parse import unquote
|
||||||
|
|
||||||
|
import wsproto
|
||||||
|
from wsproto import ConnectionType, events
|
||||||
|
from wsproto.connection import ConnectionState
|
||||||
|
from wsproto.extensions import Extension, PerMessageDeflate
|
||||||
|
from wsproto.utilities import LocalProtocolError, RemoteProtocolError
|
||||||
|
|
||||||
|
from uvicorn._types import ASGI3Application, ASGISendEvent, WebSocketEvent, WebSocketReceiveEvent, WebSocketScope
|
||||||
|
from uvicorn.config import Config
|
||||||
|
from uvicorn.logging import TRACE_LOG_LEVEL
|
||||||
|
from uvicorn.protocols.utils import (
|
||||||
|
ClientDisconnected,
|
||||||
|
get_client_addr,
|
||||||
|
get_local_addr,
|
||||||
|
get_path_with_query_string,
|
||||||
|
get_remote_addr,
|
||||||
|
is_ssl,
|
||||||
|
)
|
||||||
|
from uvicorn.server import ServerState
|
||||||
|
|
||||||
|
|
||||||
|
class FrameTooLargeError(Exception):
|
||||||
|
"""Raised when accumulated websocket message bytes exceed `ws_max_size`."""
|
||||||
|
|
||||||
|
|
||||||
|
class WebsocketBuffer:
|
||||||
|
def __init__(self, max_length: int) -> None:
|
||||||
|
self.value: BytesIO | StringIO | None = None
|
||||||
|
self.length = 0
|
||||||
|
self.max_length = max_length
|
||||||
|
|
||||||
|
def extend(self, event: events.TextMessage | events.BytesMessage) -> None:
|
||||||
|
if self.value is None:
|
||||||
|
self.value = StringIO() if isinstance(event, events.TextMessage) else BytesIO()
|
||||||
|
self.value.write(event.data) # type: ignore[arg-type]
|
||||||
|
# `ws_max_size` is a byte budget, so count UTF-8 bytes for text.
|
||||||
|
self.length += len(event.data.encode()) if isinstance(event, events.TextMessage) else len(event.data)
|
||||||
|
if self.length > self.max_length:
|
||||||
|
raise FrameTooLargeError
|
||||||
|
|
||||||
|
def clear(self) -> None:
|
||||||
|
self.value = None
|
||||||
|
self.length = 0
|
||||||
|
|
||||||
|
def to_message(self) -> WebSocketReceiveEvent:
|
||||||
|
if isinstance(self.value, StringIO):
|
||||||
|
return {"type": "websocket.receive", "text": self.value.getvalue()}
|
||||||
|
assert isinstance(self.value, BytesIO)
|
||||||
|
return {"type": "websocket.receive", "bytes": self.value.getvalue()}
|
||||||
|
|
||||||
|
|
||||||
|
class WSProtocol(asyncio.Protocol):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: Config,
|
||||||
|
server_state: ServerState,
|
||||||
|
app_state: dict[str, Any],
|
||||||
|
_loop: asyncio.AbstractEventLoop | None = None,
|
||||||
|
) -> None:
|
||||||
|
if not config.loaded:
|
||||||
|
config.load() # pragma: full coverage
|
||||||
|
|
||||||
|
self.config = config
|
||||||
|
self.app = cast(ASGI3Application, config.loaded_app)
|
||||||
|
self.loop = _loop or asyncio.get_event_loop()
|
||||||
|
self.logger = logging.getLogger("uvicorn.error")
|
||||||
|
self.root_path = config.root_path
|
||||||
|
self.app_state = app_state
|
||||||
|
|
||||||
|
# Shared server state
|
||||||
|
self.connections = server_state.connections
|
||||||
|
self.tasks = server_state.tasks
|
||||||
|
self.default_headers = server_state.default_headers
|
||||||
|
|
||||||
|
# Connection state
|
||||||
|
self.transport: asyncio.Transport = None # type: ignore[assignment]
|
||||||
|
self.server: tuple[str, int | None] | None = None
|
||||||
|
self.client: tuple[str, int] | None = None
|
||||||
|
self.scheme: Literal["wss", "ws"] = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
# WebSocket state
|
||||||
|
self.queue: asyncio.Queue[WebSocketEvent] = asyncio.Queue()
|
||||||
|
self.handshake_complete = False
|
||||||
|
self.close_sent = False
|
||||||
|
|
||||||
|
# Rejection state
|
||||||
|
self.response_started = False
|
||||||
|
|
||||||
|
self.conn = wsproto.WSConnection(connection_type=ConnectionType.SERVER)
|
||||||
|
|
||||||
|
self.read_paused = False
|
||||||
|
self.writable = asyncio.Event()
|
||||||
|
self.writable.set()
|
||||||
|
|
||||||
|
# Keepalive state
|
||||||
|
self.ping_interval = config.ws_ping_interval
|
||||||
|
self.ping_timeout = config.ws_ping_timeout
|
||||||
|
self.ping_timer: TimerHandle | None = None
|
||||||
|
self.pong_timer: TimerHandle | None = None
|
||||||
|
self.pending_ping_payload: bytes | None = None
|
||||||
|
self.ping_sent_at: float = 0.0
|
||||||
|
self.last_ping_rtt: float = 0.0
|
||||||
|
|
||||||
|
# Buffer
|
||||||
|
self.buffer = WebsocketBuffer(self.config.ws_max_size)
|
||||||
|
|
||||||
|
# Protocol interface
|
||||||
|
|
||||||
|
def connection_made(self, transport: asyncio.Transport) -> None: # type: ignore[override]
|
||||||
|
self.connections.add(self)
|
||||||
|
self.transport = transport
|
||||||
|
self.server = get_local_addr(transport)
|
||||||
|
self.client = get_remote_addr(transport)
|
||||||
|
self.scheme = "wss" if is_ssl(transport) else "ws"
|
||||||
|
|
||||||
|
if self.logger.level <= TRACE_LOG_LEVEL:
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket connection made", prefix)
|
||||||
|
|
||||||
|
def connection_lost(self, exc: Exception | None) -> None:
|
||||||
|
self.stop_keepalive()
|
||||||
|
code = 1005 if self.handshake_complete else 1006
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": code})
|
||||||
|
self.connections.remove(self)
|
||||||
|
|
||||||
|
if self.logger.level <= TRACE_LOG_LEVEL:
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket connection lost", prefix)
|
||||||
|
|
||||||
|
self.handshake_complete = True
|
||||||
|
if exc is None:
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def eof_received(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def data_received(self, data: bytes) -> None:
|
||||||
|
try:
|
||||||
|
self.conn.receive_data(data)
|
||||||
|
except RemoteProtocolError as err:
|
||||||
|
# TODO: Remove `type: ignore` when wsproto fixes the type annotation.
|
||||||
|
self.transport.write(self.conn.send(err.event_hint)) # type: ignore[arg-type] # noqa: E501
|
||||||
|
self.transport.close()
|
||||||
|
else:
|
||||||
|
self.handle_events()
|
||||||
|
|
||||||
|
def handle_events(self) -> None:
|
||||||
|
for event in self.conn.events():
|
||||||
|
if self.close_sent:
|
||||||
|
return
|
||||||
|
if isinstance(event, events.Request):
|
||||||
|
self.handle_connect(event)
|
||||||
|
elif isinstance(event, (events.TextMessage, events.BytesMessage)):
|
||||||
|
self.handle_message(event)
|
||||||
|
elif isinstance(event, events.CloseConnection):
|
||||||
|
self.handle_close(event)
|
||||||
|
elif isinstance(event, events.Ping):
|
||||||
|
self.handle_ping(event)
|
||||||
|
elif isinstance(event, events.Pong):
|
||||||
|
self.handle_pong(event)
|
||||||
|
|
||||||
|
def pause_writing(self) -> None:
|
||||||
|
"""
|
||||||
|
Called by the transport when the write buffer exceeds the high water mark.
|
||||||
|
"""
|
||||||
|
self.writable.clear() # pragma: full coverage
|
||||||
|
|
||||||
|
def resume_writing(self) -> None:
|
||||||
|
"""
|
||||||
|
Called by the transport when the write buffer drops below the low water mark.
|
||||||
|
"""
|
||||||
|
self.writable.set() # pragma: full coverage
|
||||||
|
|
||||||
|
def shutdown(self) -> None:
|
||||||
|
self.stop_keepalive()
|
||||||
|
if self.handshake_complete:
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": 1012})
|
||||||
|
output = self.conn.send(wsproto.events.CloseConnection(code=1012))
|
||||||
|
self.transport.write(output)
|
||||||
|
else:
|
||||||
|
self.send_500_response()
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def on_task_complete(self, task: asyncio.Task[None]) -> None:
|
||||||
|
self.tasks.discard(task)
|
||||||
|
|
||||||
|
# Event handlers
|
||||||
|
|
||||||
|
def handle_connect(self, event: events.Request) -> None:
|
||||||
|
headers = [(b"host", event.host.encode())]
|
||||||
|
headers += [(key.lower(), value) for key, value in event.extra_headers]
|
||||||
|
raw_path, _, query_string = event.target.partition("?")
|
||||||
|
path = unquote(raw_path)
|
||||||
|
full_path = self.root_path + path
|
||||||
|
full_raw_path = self.root_path.encode("ascii") + raw_path.encode("ascii")
|
||||||
|
self.scope: WebSocketScope = {
|
||||||
|
"type": "websocket",
|
||||||
|
"asgi": {"version": self.config.asgi_version, "spec_version": "2.4"},
|
||||||
|
"http_version": "1.1",
|
||||||
|
"scheme": self.scheme,
|
||||||
|
"server": self.server,
|
||||||
|
"client": self.client,
|
||||||
|
"root_path": self.root_path,
|
||||||
|
"path": full_path,
|
||||||
|
"raw_path": full_raw_path,
|
||||||
|
"query_string": query_string.encode("ascii"),
|
||||||
|
"headers": headers,
|
||||||
|
"subprotocols": event.subprotocols,
|
||||||
|
"state": self.app_state.copy(),
|
||||||
|
"extensions": {"websocket.http.response": {}},
|
||||||
|
}
|
||||||
|
self.queue.put_nowait({"type": "websocket.connect"})
|
||||||
|
task = self.loop.create_task(self.run_asgi())
|
||||||
|
task.add_done_callback(self.on_task_complete)
|
||||||
|
self.tasks.add(task)
|
||||||
|
|
||||||
|
def handle_message(self, event: events.TextMessage | events.BytesMessage) -> None:
|
||||||
|
try:
|
||||||
|
self.buffer.extend(event)
|
||||||
|
except FrameTooLargeError:
|
||||||
|
self.close_sent = True
|
||||||
|
reason = f"Message exceeds the maximum size ({self.config.ws_max_size} bytes)"
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": 1009, "reason": reason})
|
||||||
|
if not self.transport.is_closing():
|
||||||
|
self.transport.write(self.conn.send(wsproto.events.CloseConnection(code=1009, reason=reason)))
|
||||||
|
self.transport.close()
|
||||||
|
return
|
||||||
|
if event.message_finished:
|
||||||
|
self.queue.put_nowait(self.buffer.to_message())
|
||||||
|
self.buffer.clear()
|
||||||
|
if not self.read_paused:
|
||||||
|
self.read_paused = True
|
||||||
|
self.transport.pause_reading()
|
||||||
|
|
||||||
|
def handle_close(self, event: events.CloseConnection) -> None:
|
||||||
|
if self.conn.state == ConnectionState.REMOTE_CLOSING:
|
||||||
|
self.transport.write(self.conn.send(event.response()))
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": event.code, "reason": event.reason})
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def handle_ping(self, event: events.Ping) -> None:
|
||||||
|
self.transport.write(self.conn.send(event.response()))
|
||||||
|
|
||||||
|
def handle_pong(self, event: events.Pong) -> None:
|
||||||
|
# Ignore unsolicited pongs and stale pongs whose payload doesn't match the ping currently in flight.
|
||||||
|
if self.pending_ping_payload is None or bytes(event.payload) != self.pending_ping_payload:
|
||||||
|
return # pragma: no cover
|
||||||
|
|
||||||
|
self.last_ping_rtt = self.loop.time() - self.ping_sent_at
|
||||||
|
self.pending_ping_payload = None
|
||||||
|
# The peer answered in time; cancel the pong deadline and chain the next ping. This `schedule_ping()` call is
|
||||||
|
# what keeps the keepalive loop running when ping_timeout is set. When ping_timeout is None the next ping is
|
||||||
|
# already scheduled by `send_keepalive_ping`, so we must not schedule a duplicate here.
|
||||||
|
if self.pong_timer is not None:
|
||||||
|
self.pong_timer.cancel()
|
||||||
|
self.pong_timer = None
|
||||||
|
self.schedule_ping()
|
||||||
|
|
||||||
|
def start_keepalive(self) -> None:
|
||||||
|
if self.ping_interval is not None and self.ping_interval > 0:
|
||||||
|
self.schedule_ping()
|
||||||
|
|
||||||
|
def stop_keepalive(self) -> None:
|
||||||
|
if self.ping_timer is not None:
|
||||||
|
self.ping_timer.cancel()
|
||||||
|
self.ping_timer = None
|
||||||
|
if self.pong_timer is not None: # pragma: no cover
|
||||||
|
self.pong_timer.cancel()
|
||||||
|
self.pong_timer = None
|
||||||
|
self.pending_ping_payload = None
|
||||||
|
|
||||||
|
def schedule_ping(self) -> None:
|
||||||
|
assert self.ping_interval is not None
|
||||||
|
delay = max(0.0, self.ping_interval - self.last_ping_rtt)
|
||||||
|
self.ping_timer = self.loop.call_later(delay, self.send_keepalive_ping)
|
||||||
|
|
||||||
|
def send_keepalive_ping(self) -> None:
|
||||||
|
self.ping_timer = None
|
||||||
|
if self.close_sent or self.transport.is_closing(): # pragma: no cover
|
||||||
|
return
|
||||||
|
# Random 4-byte payload identifies this ping; `handle_pong` uses it to ignore stale or unsolicited pongs.
|
||||||
|
self.pending_ping_payload = struct.pack("!I", random.getrandbits(32))
|
||||||
|
self.ping_sent_at = self.loop.time()
|
||||||
|
self.transport.write(self.conn.send(wsproto.events.Ping(payload=self.pending_ping_payload)))
|
||||||
|
if self.ping_timeout is not None:
|
||||||
|
self.pong_timer = self.loop.call_later(self.ping_timeout, self.keepalive_timeout)
|
||||||
|
else: # pragma: no cover
|
||||||
|
self.schedule_ping()
|
||||||
|
|
||||||
|
def keepalive_timeout(self) -> None:
|
||||||
|
self.pong_timer = None
|
||||||
|
self.pending_ping_payload = None
|
||||||
|
if self.close_sent or self.transport.is_closing(): # pragma: no cover
|
||||||
|
return
|
||||||
|
if self.logger.level <= TRACE_LOG_LEVEL:
|
||||||
|
prefix = "%s:%d - " % self.client if self.client else ""
|
||||||
|
self.logger.log(TRACE_LOG_LEVEL, "%sWebSocket keepalive ping timeout", prefix)
|
||||||
|
reason = "keepalive ping timeout"
|
||||||
|
self.transport.write(self.conn.send(wsproto.events.CloseConnection(code=1011, reason=reason)))
|
||||||
|
self.close_sent = True
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
def send_500_response(self) -> None:
|
||||||
|
if self.response_started or self.handshake_complete:
|
||||||
|
return # we cannot send responses anymore
|
||||||
|
headers: list[tuple[bytes, bytes]] = [
|
||||||
|
(b"content-type", b"text/plain; charset=utf-8"),
|
||||||
|
(b"connection", b"close"),
|
||||||
|
(b"content-length", b"21"),
|
||||||
|
]
|
||||||
|
output = self.conn.send(wsproto.events.RejectConnection(status_code=500, headers=headers, has_body=True))
|
||||||
|
output += self.conn.send(wsproto.events.RejectData(data=b"Internal Server Error"))
|
||||||
|
self.transport.write(output)
|
||||||
|
|
||||||
|
async def run_asgi(self) -> None:
|
||||||
|
try:
|
||||||
|
result = await self.app(self.scope, self.receive, self.send) # type: ignore[func-returns-value]
|
||||||
|
except ClientDisconnected:
|
||||||
|
pass # pragma: full coverage
|
||||||
|
except BaseException:
|
||||||
|
self.logger.exception("Exception in ASGI application\n")
|
||||||
|
self.send_500_response()
|
||||||
|
else:
|
||||||
|
if not self.handshake_complete:
|
||||||
|
self.logger.error("ASGI callable returned without completing handshake.")
|
||||||
|
self.send_500_response()
|
||||||
|
elif result is not None:
|
||||||
|
self.logger.error("ASGI callable should return None, but returned '%s'.", result)
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
async def send(self, message: ASGISendEvent) -> None:
|
||||||
|
await self.writable.wait()
|
||||||
|
|
||||||
|
if not self.handshake_complete:
|
||||||
|
if message["type"] == "websocket.accept":
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" [accepted]',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
)
|
||||||
|
subprotocol = message.get("subprotocol")
|
||||||
|
extra_headers = self.default_headers + list(message.get("headers", []))
|
||||||
|
extensions: list[Extension] = []
|
||||||
|
if self.config.ws_per_message_deflate:
|
||||||
|
extensions.append(PerMessageDeflate())
|
||||||
|
if not self.transport.is_closing():
|
||||||
|
self.handshake_complete = True
|
||||||
|
output = self.conn.send(
|
||||||
|
wsproto.events.AcceptConnection(
|
||||||
|
subprotocol=subprotocol,
|
||||||
|
extensions=extensions,
|
||||||
|
extra_headers=extra_headers,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.transport.write(output)
|
||||||
|
self.start_keepalive()
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.close":
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": 1006})
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" 403',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
)
|
||||||
|
self.handshake_complete = True
|
||||||
|
self.close_sent = True
|
||||||
|
event = events.RejectConnection(status_code=403, headers=[])
|
||||||
|
output = self.conn.send(event)
|
||||||
|
self.transport.write(output)
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.http.response.start":
|
||||||
|
# ensure status code is in the valid range
|
||||||
|
if not (100 <= message["status"] < 600):
|
||||||
|
msg = "Invalid HTTP status code '%d' in response."
|
||||||
|
raise RuntimeError(msg % message["status"])
|
||||||
|
self.logger.info(
|
||||||
|
'%s - "WebSocket %s" %d',
|
||||||
|
get_client_addr(self.scope),
|
||||||
|
get_path_with_query_string(self.scope),
|
||||||
|
message["status"],
|
||||||
|
)
|
||||||
|
self.handshake_complete = True
|
||||||
|
event = events.RejectConnection(
|
||||||
|
status_code=message["status"],
|
||||||
|
headers=list(message["headers"]),
|
||||||
|
has_body=True,
|
||||||
|
)
|
||||||
|
output = self.conn.send(event)
|
||||||
|
self.transport.write(output)
|
||||||
|
self.response_started = True
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Expected ASGI message 'websocket.accept', 'websocket.close' "
|
||||||
|
f"or 'websocket.http.response.start' but got '{message['type']}'."
|
||||||
|
)
|
||||||
|
|
||||||
|
elif not self.close_sent and not self.response_started:
|
||||||
|
try:
|
||||||
|
if message["type"] == "websocket.send":
|
||||||
|
bytes_data = message.get("bytes")
|
||||||
|
text_data = message.get("text")
|
||||||
|
data = text_data if bytes_data is None else bytes_data
|
||||||
|
output = self.conn.send(wsproto.events.Message(data=data)) # type: ignore
|
||||||
|
if not self.transport.is_closing():
|
||||||
|
self.transport.write(output)
|
||||||
|
|
||||||
|
elif message["type"] == "websocket.close":
|
||||||
|
self.close_sent = True
|
||||||
|
code = message.get("code", 1000)
|
||||||
|
reason = message.get("reason", "") or ""
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": code, "reason": reason})
|
||||||
|
output = self.conn.send(wsproto.events.CloseConnection(code=code, reason=reason))
|
||||||
|
if not self.transport.is_closing():
|
||||||
|
self.transport.write(output)
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Expected ASGI message 'websocket.send' or 'websocket.close', but got '{message['type']}'."
|
||||||
|
)
|
||||||
|
except LocalProtocolError as exc:
|
||||||
|
raise ClientDisconnected from exc
|
||||||
|
elif self.response_started:
|
||||||
|
if message["type"] == "websocket.http.response.body":
|
||||||
|
body_finished = not message.get("more_body", False)
|
||||||
|
reject_data = events.RejectData(data=message["body"], body_finished=body_finished)
|
||||||
|
output = self.conn.send(reject_data)
|
||||||
|
self.transport.write(output)
|
||||||
|
|
||||||
|
if body_finished:
|
||||||
|
self.queue.put_nowait({"type": "websocket.disconnect", "code": 1006})
|
||||||
|
self.close_sent = True
|
||||||
|
self.transport.close()
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(f"Expected ASGI message 'websocket.http.response.body' but got '{message['type']}'.")
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise RuntimeError(f"Unexpected ASGI message '{message['type']}', after sending 'websocket.close'.")
|
||||||
|
|
||||||
|
async def receive(self) -> WebSocketEvent:
|
||||||
|
message = await self.queue.get()
|
||||||
|
if self.read_paused and self.queue.empty():
|
||||||
|
self.read_paused = False
|
||||||
|
self.transport.resume_reading()
|
||||||
|
return message
|
||||||
Reference in New Issue
Block a user