Coverage for /pythoncovmergedfiles/medio/medio/src/aiohttp/aiohttp/web_ws.py: 20%
320 statements
« prev ^ index » next coverage.py v7.3.1, created at 2023-09-27 06:09 +0000
« prev ^ index » next coverage.py v7.3.1, created at 2023-09-27 06:09 +0000
1import asyncio
2import base64
3import binascii
4import dataclasses
5import hashlib
6import json
7import sys
8from typing import Any, Final, Iterable, Optional, Tuple, cast
10from multidict import CIMultiDict
12from . import hdrs
13from .abc import AbstractStreamWriter
14from .helpers import call_later, set_result
15from .http import (
16 WS_CLOSED_MESSAGE,
17 WS_CLOSING_MESSAGE,
18 WS_KEY,
19 WebSocketError,
20 WebSocketReader,
21 WebSocketWriter,
22 WSCloseCode,
23 WSMessage,
24 WSMsgType as WSMsgType,
25 ws_ext_gen,
26 ws_ext_parse,
27)
28from .log import ws_logger
29from .streams import EofStream, FlowControlDataQueue
30from .typedefs import JSONDecoder, JSONEncoder
31from .web_exceptions import HTTPBadRequest, HTTPException
32from .web_request import BaseRequest
33from .web_response import StreamResponse
35if sys.version_info >= (3, 11):
36 import asyncio as async_timeout
37else:
38 import async_timeout
40__all__ = (
41 "WebSocketResponse",
42 "WebSocketReady",
43 "WSMsgType",
44)
46THRESHOLD_CONNLOST_ACCESS: Final[int] = 5
49@dataclasses.dataclass(frozen=True)
50class WebSocketReady:
51 ok: bool
52 protocol: Optional[str]
54 def __bool__(self) -> bool:
55 return self.ok
58class WebSocketResponse(StreamResponse):
59 __slots__ = (
60 "_protocols",
61 "_ws_protocol",
62 "_writer",
63 "_reader",
64 "_closed",
65 "_closing",
66 "_conn_lost",
67 "_close_code",
68 "_loop",
69 "_waiting",
70 "_exception",
71 "_timeout",
72 "_receive_timeout",
73 "_autoclose",
74 "_autoping",
75 "_heartbeat",
76 "_heartbeat_cb",
77 "_pong_heartbeat",
78 "_pong_response_cb",
79 "_compress",
80 "_max_msg_size",
81 )
83 def __init__(
84 self,
85 *,
86 timeout: float = 10.0,
87 receive_timeout: Optional[float] = None,
88 autoclose: bool = True,
89 autoping: bool = True,
90 heartbeat: Optional[float] = None,
91 protocols: Iterable[str] = (),
92 compress: bool = True,
93 max_msg_size: int = 4 * 1024 * 1024,
94 ) -> None:
95 super().__init__(status=101)
96 self._length_check = False
97 self._protocols = protocols
98 self._ws_protocol: Optional[str] = None
99 self._writer: Optional[WebSocketWriter] = None
100 self._reader: Optional[FlowControlDataQueue[WSMessage]] = None
101 self._closed = False
102 self._closing = False
103 self._conn_lost = 0
104 self._close_code: Optional[int] = None
105 self._loop: Optional[asyncio.AbstractEventLoop] = None
106 self._waiting: Optional[asyncio.Future[bool]] = None
107 self._exception: Optional[BaseException] = None
108 self._timeout = timeout
109 self._receive_timeout = receive_timeout
110 self._autoclose = autoclose
111 self._autoping = autoping
112 self._heartbeat = heartbeat
113 self._heartbeat_cb: Optional[asyncio.TimerHandle] = None
114 if heartbeat is not None:
115 self._pong_heartbeat = heartbeat / 2.0
116 self._pong_response_cb: Optional[asyncio.TimerHandle] = None
117 self._compress = compress
118 self._max_msg_size = max_msg_size
120 def _cancel_heartbeat(self) -> None:
121 if self._pong_response_cb is not None:
122 self._pong_response_cb.cancel()
123 self._pong_response_cb = None
125 if self._heartbeat_cb is not None:
126 self._heartbeat_cb.cancel()
127 self._heartbeat_cb = None
129 def _reset_heartbeat(self) -> None:
130 self._cancel_heartbeat()
132 if self._heartbeat is not None:
133 assert self._loop is not None
134 self._heartbeat_cb = call_later(
135 self._send_heartbeat,
136 self._heartbeat,
137 self._loop,
138 timeout_ceil_threshold=self._req._protocol._timeout_ceil_threshold
139 if self._req is not None
140 else 5,
141 )
143 def _send_heartbeat(self) -> None:
144 if self._heartbeat is not None and not self._closed:
145 assert self._loop is not None and self._writer is not None
146 # fire-and-forget a task is not perfect but maybe ok for
147 # sending ping. Otherwise we need a long-living heartbeat
148 # task in the class.
149 self._loop.create_task(self._writer.ping()) # type: ignore[unused-awaitable]
151 if self._pong_response_cb is not None:
152 self._pong_response_cb.cancel()
153 self._pong_response_cb = call_later(
154 self._pong_not_received,
155 self._pong_heartbeat,
156 self._loop,
157 timeout_ceil_threshold=self._req._protocol._timeout_ceil_threshold
158 if self._req is not None
159 else 5,
160 )
162 def _pong_not_received(self) -> None:
163 if self._req is not None and self._req.transport is not None:
164 self._closed = True
165 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
166 self._exception = asyncio.TimeoutError()
167 self._req.transport.close()
169 async def prepare(self, request: BaseRequest) -> AbstractStreamWriter:
170 # make pre-check to don't hide it by do_handshake() exceptions
171 if self._payload_writer is not None:
172 return self._payload_writer
174 protocol, writer = self._pre_start(request)
175 payload_writer = await super().prepare(request)
176 assert payload_writer is not None
177 self._post_start(request, protocol, writer)
178 await payload_writer.drain()
179 return payload_writer
181 def _handshake(
182 self, request: BaseRequest
183 ) -> Tuple["CIMultiDict[str]", str, bool, bool]:
184 headers = request.headers
185 if "websocket" != headers.get(hdrs.UPGRADE, "").lower().strip():
186 raise HTTPBadRequest(
187 text=(
188 "No WebSocket UPGRADE hdr: {}\n Can "
189 '"Upgrade" only to "WebSocket".'
190 ).format(headers.get(hdrs.UPGRADE))
191 )
193 if "upgrade" not in headers.get(hdrs.CONNECTION, "").lower():
194 raise HTTPBadRequest(
195 text="No CONNECTION upgrade hdr: {}".format(
196 headers.get(hdrs.CONNECTION)
197 )
198 )
200 # find common sub-protocol between client and server
201 protocol = None
202 if hdrs.SEC_WEBSOCKET_PROTOCOL in headers:
203 req_protocols = [
204 str(proto.strip())
205 for proto in headers[hdrs.SEC_WEBSOCKET_PROTOCOL].split(",")
206 ]
208 for proto in req_protocols:
209 if proto in self._protocols:
210 protocol = proto
211 break
212 else:
213 # No overlap found: Return no protocol as per spec
214 ws_logger.warning(
215 "Client protocols %r don’t overlap server-known ones %r",
216 req_protocols,
217 self._protocols,
218 )
220 # check supported version
221 version = headers.get(hdrs.SEC_WEBSOCKET_VERSION, "")
222 if version not in ("13", "8", "7"):
223 raise HTTPBadRequest(text=f"Unsupported version: {version}")
225 # check client handshake for validity
226 key = headers.get(hdrs.SEC_WEBSOCKET_KEY)
227 try:
228 if not key or len(base64.b64decode(key)) != 16:
229 raise HTTPBadRequest(text=f"Handshake error: {key!r}")
230 except binascii.Error:
231 raise HTTPBadRequest(text=f"Handshake error: {key!r}") from None
233 accept_val = base64.b64encode(
234 hashlib.sha1(key.encode() + WS_KEY).digest()
235 ).decode()
236 response_headers = CIMultiDict(
237 {
238 hdrs.UPGRADE: "websocket",
239 hdrs.CONNECTION: "upgrade",
240 hdrs.SEC_WEBSOCKET_ACCEPT: accept_val,
241 }
242 )
244 notakeover = False
245 compress = 0
246 if self._compress:
247 extensions = headers.get(hdrs.SEC_WEBSOCKET_EXTENSIONS)
248 # Server side always get return with no exception.
249 # If something happened, just drop compress extension
250 compress, notakeover = ws_ext_parse(extensions, isserver=True)
251 if compress:
252 enabledext = ws_ext_gen(
253 compress=compress, isserver=True, server_notakeover=notakeover
254 )
255 response_headers[hdrs.SEC_WEBSOCKET_EXTENSIONS] = enabledext
257 if protocol:
258 response_headers[hdrs.SEC_WEBSOCKET_PROTOCOL] = protocol
259 return (
260 response_headers,
261 protocol,
262 compress,
263 notakeover,
264 ) # type: ignore[return-value]
266 def _pre_start(self, request: BaseRequest) -> Tuple[str, WebSocketWriter]:
267 self._loop = request._loop
269 headers, protocol, compress, notakeover = self._handshake(request)
271 self.set_status(101)
272 self.headers.update(headers)
273 self.force_close()
274 self._compress = compress
275 transport = request._protocol.transport
276 assert transport is not None
277 writer = WebSocketWriter(
278 request._protocol, transport, compress=compress, notakeover=notakeover
279 )
281 return protocol, writer
283 def _post_start(
284 self, request: BaseRequest, protocol: str, writer: WebSocketWriter
285 ) -> None:
286 self._ws_protocol = protocol
287 self._writer = writer
289 self._reset_heartbeat()
291 loop = self._loop
292 assert loop is not None
293 self._reader = FlowControlDataQueue(request._protocol, 2**16, loop=loop)
294 request.protocol.set_parser(
295 WebSocketReader(self._reader, self._max_msg_size, compress=self._compress)
296 )
297 # disable HTTP keepalive for WebSocket
298 request.protocol.keep_alive(False)
300 def can_prepare(self, request: BaseRequest) -> WebSocketReady:
301 if self._writer is not None:
302 raise RuntimeError("Already started")
303 try:
304 _, protocol, _, _ = self._handshake(request)
305 except HTTPException:
306 return WebSocketReady(False, None)
307 else:
308 return WebSocketReady(True, protocol)
310 @property
311 def closed(self) -> bool:
312 return self._closed
314 @property
315 def close_code(self) -> Optional[int]:
316 return self._close_code
318 @property
319 def ws_protocol(self) -> Optional[str]:
320 return self._ws_protocol
322 @property
323 def compress(self) -> bool:
324 return self._compress
326 def exception(self) -> Optional[BaseException]:
327 return self._exception
329 async def ping(self, message: bytes = b"") -> None:
330 if self._writer is None:
331 raise RuntimeError("Call .prepare() first")
332 await self._writer.ping(message)
334 async def pong(self, message: bytes = b"") -> None:
335 # unsolicited pong
336 if self._writer is None:
337 raise RuntimeError("Call .prepare() first")
338 await self._writer.pong(message)
340 async def send_str(self, data: str, compress: Optional[bool] = None) -> None:
341 if self._writer is None:
342 raise RuntimeError("Call .prepare() first")
343 if not isinstance(data, str):
344 raise TypeError("data argument must be str (%r)" % type(data))
345 await self._writer.send(data, binary=False, compress=compress)
347 async def send_bytes(self, data: bytes, compress: Optional[bool] = None) -> None:
348 if self._writer is None:
349 raise RuntimeError("Call .prepare() first")
350 if not isinstance(data, (bytes, bytearray, memoryview)):
351 raise TypeError("data argument must be byte-ish (%r)" % type(data))
352 await self._writer.send(data, binary=True, compress=compress)
354 async def send_json(
355 self,
356 data: Any,
357 compress: Optional[bool] = None,
358 *,
359 dumps: JSONEncoder = json.dumps,
360 ) -> None:
361 await self.send_str(dumps(data), compress=compress)
363 async def write_eof(self) -> None: # type: ignore[override]
364 if self._eof_sent:
365 return
366 if self._payload_writer is None:
367 raise RuntimeError("Response has not been started")
369 await self.close()
370 self._eof_sent = True
372 async def close(self, *, code: int = WSCloseCode.OK, message: bytes = b"") -> bool:
373 if self._writer is None:
374 raise RuntimeError("Call .prepare() first")
376 self._cancel_heartbeat()
377 reader = self._reader
378 assert reader is not None
380 # we need to break `receive()` cycle first,
381 # `close()` may be called from different task
382 if self._waiting is not None and not self._closed:
383 reader.feed_data(WS_CLOSING_MESSAGE, 0)
384 await self._waiting
386 if not self._closed:
387 self._closed = True
388 try:
389 await self._writer.close(code, message)
390 writer = self._payload_writer
391 assert writer is not None
392 await writer.drain()
393 except (asyncio.CancelledError, asyncio.TimeoutError):
394 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
395 raise
396 except Exception as exc:
397 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
398 self._exception = exc
399 return True
401 if self._closing:
402 return True
404 reader = self._reader
405 assert reader is not None
406 try:
407 async with async_timeout.timeout(self._timeout):
408 msg = await reader.read()
409 except asyncio.CancelledError:
410 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
411 raise
412 except Exception as exc:
413 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
414 self._exception = exc
415 return True
417 if msg.type == WSMsgType.CLOSE:
418 self._close_code = msg.data
419 return True
421 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
422 self._exception = asyncio.TimeoutError()
423 return True
424 else:
425 return False
427 async def receive(self, timeout: Optional[float] = None) -> WSMessage:
428 if self._reader is None:
429 raise RuntimeError("Call .prepare() first")
431 loop = self._loop
432 assert loop is not None
433 while True:
434 if self._waiting is not None:
435 raise RuntimeError("Concurrent call to receive() is not allowed")
437 if self._closed:
438 self._conn_lost += 1
439 if self._conn_lost >= THRESHOLD_CONNLOST_ACCESS:
440 raise RuntimeError("WebSocket connection is closed.")
441 return WS_CLOSED_MESSAGE
442 elif self._closing:
443 return WS_CLOSING_MESSAGE
445 try:
446 self._waiting = loop.create_future()
447 try:
448 async with async_timeout.timeout(timeout or self._receive_timeout):
449 msg = await self._reader.read()
450 self._reset_heartbeat()
451 finally:
452 waiter = self._waiting
453 set_result(waiter, True)
454 self._waiting = None
455 except (asyncio.CancelledError, asyncio.TimeoutError):
456 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
457 raise
458 except EofStream:
459 self._close_code = WSCloseCode.OK
460 await self.close()
461 return WSMessage(WSMsgType.CLOSED, None, None)
462 except WebSocketError as exc:
463 self._close_code = exc.code
464 await self.close(code=exc.code)
465 return WSMessage(WSMsgType.ERROR, exc, None)
466 except Exception as exc:
467 self._exception = exc
468 self._closing = True
469 self._close_code = WSCloseCode.ABNORMAL_CLOSURE
470 await self.close()
471 return WSMessage(WSMsgType.ERROR, exc, None)
473 if msg.type == WSMsgType.CLOSE:
474 self._closing = True
475 self._close_code = msg.data
476 # Could be closed while awaiting reader.
477 if not self._closed and self._autoclose: # type: ignore[redundant-expr]
478 await self.close()
479 elif msg.type == WSMsgType.CLOSING:
480 self._closing = True
481 elif msg.type == WSMsgType.PING and self._autoping:
482 await self.pong(msg.data)
483 continue
484 elif msg.type == WSMsgType.PONG and self._autoping:
485 continue
487 return msg
489 async def receive_str(self, *, timeout: Optional[float] = None) -> str:
490 msg = await self.receive(timeout)
491 if msg.type != WSMsgType.TEXT:
492 raise TypeError(
493 "Received message {}:{!r} is not WSMsgType.TEXT".format(
494 msg.type, msg.data
495 )
496 )
497 return cast(str, msg.data)
499 async def receive_bytes(self, *, timeout: Optional[float] = None) -> bytes:
500 msg = await self.receive(timeout)
501 if msg.type != WSMsgType.BINARY:
502 raise TypeError(f"Received message {msg.type}:{msg.data!r} is not bytes")
503 return cast(bytes, msg.data)
505 async def receive_json(
506 self, *, loads: JSONDecoder = json.loads, timeout: Optional[float] = None
507 ) -> Any:
508 data = await self.receive_str(timeout=timeout)
509 return loads(data)
511 async def write(self, data: bytes) -> None:
512 raise RuntimeError("Cannot call .write() for websocket")
514 def __aiter__(self) -> "WebSocketResponse":
515 return self
517 async def __anext__(self) -> WSMessage:
518 msg = await self.receive()
519 if msg.type in (WSMsgType.CLOSE, WSMsgType.CLOSING, WSMsgType.CLOSED):
520 raise StopAsyncIteration
521 return msg
523 def _cancel(self, exc: BaseException) -> None:
524 if self._reader is not None:
525 self._reader.set_exception(exc)