codekingpro/portable-devtools
114k
1"""2wsproto/handshake3~~~~~~~~~~~~~~~~~~4 5An implementation of WebSocket handshakes.6"""7from __future__ import annotations8 9from collections import deque10from typing import (11 TYPE_CHECKING,12 cast,13)14 15import h1116 17from .connection import Connection, ConnectionState, ConnectionType18from .events import AcceptConnection, Event, RejectConnection, RejectData, Request19from .extensions import Extension20from .utilities import (21 LocalProtocolError,22 RemoteProtocolError,23 generate_accept_token,24 generate_nonce,25 normed_header_dict,26 split_comma_header,27)28 29if TYPE_CHECKING:30 from collections.abc import Generator, Iterable, Sequence31 32 from .typing import Headers33 34# RFC6455, Section 4.2.1/6 - Reading the Client's Opening Handshake35WEBSOCKET_VERSION = b"13"36 37# RFC6455, Section 4.2.1/3 - Value of the Upgrade header38WEBSOCKET_UPGRADE = b"websocket"39 40 41class H11Handshake:42 """A Handshake implementation for HTTP/1.1 connections."""43 44 def __init__(self, connection_type: ConnectionType) -> None:45 self.client = connection_type is ConnectionType.CLIENT46 self._state = ConnectionState.CONNECTING47 48 if self.client:49 self._h11_connection = h11.Connection(h11.CLIENT)50 else:51 self._h11_connection = h11.Connection(h11.SERVER)52 53 self._connection: Connection | None = None54 self._events: deque[Event] = deque()55 self._initiating_request: Request | None = None56 self._nonce: bytes | None = None57 58 @property59 def state(self) -> ConnectionState:60 return self._state61 62 @property63 def connection(self) -> Connection | None:64 """65 Return the established connection.66 67 This will either return the connection or raise a68 LocalProtocolError if the connection has not yet been69 established.70 71 :rtype: h11.Connection72 """73 return self._connection74 75 def initiate_upgrade_connection(76 self, headers: Headers, path: bytes | str,77 ) -> None:78 """79 Initiate an upgrade connection.80 81 This should be used if the request has already be received and82 parsed.83 84 :param list headers: HTTP headers represented as a list of 2-tuples.85 :param str path: A URL path.86 """87 if self.client:88 msg = "Cannot initiate an upgrade connection when acting as the client"89 raise LocalProtocolError(90 msg,91 )92 upgrade_request = h11.Request(method=b"GET", target=path, headers=headers)93 h11_client = h11.Connection(h11.CLIENT)94 self.receive_data(h11_client.send(upgrade_request))95 96 def send(self, event: Event) -> bytes:97 """98 Send an event to the remote.99 100 This will return the bytes to send based on the event or raise101 a LocalProtocolError if the event is not valid given the102 state.103 104 :returns: Data to send to the WebSocket peer.105 :rtype: bytes106 """107 data = b""108 if isinstance(event, Request):109 data += self._initiate_connection(event)110 elif isinstance(event, AcceptConnection):111 data += self._accept(event)112 elif isinstance(event, RejectConnection):113 data += self._reject(event)114 elif isinstance(event, RejectData):115 data += self._send_reject_data(event)116 else:117 msg = f"Event {event} cannot be sent during the handshake"118 raise LocalProtocolError(119 msg,120 )121 return data122 123 def receive_data(self, data: bytes | None) -> None:124 """125 Receive data from the remote.126 127 A list of events that the remote peer triggered by sending128 this data can be retrieved with :meth:`events`.129 130 :param bytes data: Data received from the WebSocket peer.131 """132 self._h11_connection.receive_data(data or b"")133 while True:134 try:135 event = self._h11_connection.next_event()136 except h11.RemoteProtocolError:137 msg = "Bad HTTP message"138 raise RemoteProtocolError(139 msg, event_hint=RejectConnection(),140 )141 if (142 isinstance(event, h11.ConnectionClosed)143 or event is h11.NEED_DATA144 or event is h11.PAUSED145 ):146 break147 148 if self.client:149 if isinstance(event, h11.InformationalResponse):150 if event.status_code == 101:151 self._events.append(self._establish_client_connection(event))152 else:153 self._events.append(154 RejectConnection(155 headers=list(event.headers),156 status_code=event.status_code,157 has_body=False,158 ),159 )160 self._state = ConnectionState.CLOSED161 elif isinstance(event, h11.Response):162 self._state = ConnectionState.REJECTING163 self._events.append(164 RejectConnection(165 headers=list(event.headers),166 status_code=event.status_code,167 has_body=True,168 ),169 )170 elif isinstance(event, h11.Data):171 self._events.append(172 RejectData(data=event.data, body_finished=False),173 )174 elif isinstance(event, h11.EndOfMessage):175 self._events.append(RejectData(data=b"", body_finished=True))176 self._state = ConnectionState.CLOSED177 elif isinstance(event, h11.Request):178 self._events.append(self._process_connection_request(event))179 180 def events(self) -> Generator[Event, None, None]:181 """182 Return a generator that provides any events that have been generated183 by protocol activity.184 185 :returns: a generator that yields H11 events.186 """187 while self._events:188 yield self._events.popleft()189 190 # Server mode methods191 192 def _process_connection_request(193 self, event: h11.Request,194 ) -> Request:195 if event.method != b"GET":196 msg = "Request method must be GET"197 raise RemoteProtocolError(198 msg, event_hint=RejectConnection(),199 )200 connection_tokens = None201 extensions: list[str] = []202 host = None203 key = None204 subprotocols: list[str] = []205 upgrade = b""206 version = None207 headers: Headers = []208 for name, value in event.headers:209 name = name.lower()210 if name == b"connection":211 connection_tokens = split_comma_header(value)212 elif name == b"host":213 host = value.decode("idna")214 continue # Skip appending to headers215 elif name == b"sec-websocket-extensions":216 extensions.extend(split_comma_header(value))217 continue # Skip appending to headers218 elif name == b"sec-websocket-key":219 key = value220 elif name == b"sec-websocket-protocol":221 subprotocols.extend(split_comma_header(value))222 continue # Skip appending to headers223 elif name == b"sec-websocket-version":224 version = value225 elif name == b"upgrade":226 upgrade = value227 headers.append((name, value))228 if connection_tokens is None or not any(229 token.lower() == "upgrade" for token in connection_tokens230 ):231 msg = "Missing header, 'Connection: Upgrade'"232 raise RemoteProtocolError(233 msg, event_hint=RejectConnection(),234 )235 if version != WEBSOCKET_VERSION:236 msg = "Missing header, 'Sec-WebSocket-Version'"237 raise RemoteProtocolError(238 msg,239 event_hint=RejectConnection(240 headers=[(b"Sec-WebSocket-Version", WEBSOCKET_VERSION)],241 status_code=426 if version else 400,242 ),243 )244 if key is None:245 msg = "Missing header, 'Sec-WebSocket-Key'"246 raise RemoteProtocolError(247 msg, event_hint=RejectConnection(),248 )249 if upgrade.lower() != WEBSOCKET_UPGRADE:250 msg = f"Missing header, 'Upgrade: {WEBSOCKET_UPGRADE.decode()}'"251 raise RemoteProtocolError(252 msg,253 event_hint=RejectConnection(),254 )255 if host is None:256 msg = "Missing header, 'Host'"257 raise RemoteProtocolError(258 msg, event_hint=RejectConnection(),259 )260 261 self._initiating_request = Request(262 extensions=extensions,263 extra_headers=headers,264 host=host,265 subprotocols=subprotocols,266 target=event.target.decode("ascii"),267 )268 return self._initiating_request269 270 def _accept(self, event: AcceptConnection) -> bytes:271 # _accept is always called after _process_connection_request.272 assert self._initiating_request is not None273 request_headers = normed_header_dict(self._initiating_request.extra_headers)274 275 nonce = request_headers[b"sec-websocket-key"]276 accept_token = generate_accept_token(nonce)277 278 headers = [279 (b"Upgrade", WEBSOCKET_UPGRADE),280 (b"Connection", b"Upgrade"),281 (b"Sec-WebSocket-Accept", accept_token),282 ]283 284 if event.subprotocol is not None:285 if event.subprotocol not in self._initiating_request.subprotocols:286 msg = f"unexpected subprotocol {event.subprotocol}"287 raise LocalProtocolError(msg)288 headers.append(289 (b"Sec-WebSocket-Protocol", event.subprotocol.encode("ascii")),290 )291 292 if event.extensions:293 accepts = server_extensions_handshake(294 cast("Sequence[str]", self._initiating_request.extensions),295 event.extensions,296 )297 if accepts:298 headers.append((b"Sec-WebSocket-Extensions", accepts))299 300 response = h11.InformationalResponse(301 status_code=101,302 headers=headers + event.extra_headers,303 reason=b"Switching Protocols",304 )305 self._connection = Connection(306 ConnectionType.CLIENT if self.client else ConnectionType.SERVER,307 event.extensions,308 )309 self._state = ConnectionState.OPEN310 return self._h11_connection.send(response) or b""311 312 def _reject(self, event: RejectConnection) -> bytes:313 if self.state != ConnectionState.CONNECTING:314 msg = f"Connection cannot be rejected in state {self.state}"315 raise LocalProtocolError(316 msg,317 )318 319 headers = list(event.headers)320 if not event.has_body:321 headers.append((b"content-length", b"0"))322 response = h11.Response(status_code=event.status_code, headers=headers)323 data = self._h11_connection.send(response) or b""324 self._state = ConnectionState.REJECTING325 if not event.has_body:326 data += self._h11_connection.send(h11.EndOfMessage()) or b""327 self._state = ConnectionState.CLOSED328 return data329 330 def _send_reject_data(self, event: RejectData) -> bytes:331 if self.state != ConnectionState.REJECTING:332 msg = f"Cannot send rejection data in state {self.state}"333 raise LocalProtocolError(334 msg,335 )336 337 data = self._h11_connection.send(h11.Data(data=event.data)) or b""338 if event.body_finished:339 data += self._h11_connection.send(h11.EndOfMessage()) or b""340 self._state = ConnectionState.CLOSED341 return data342 343 # Client mode methods344 345 def _initiate_connection(self, request: Request) -> bytes:346 self._initiating_request = request347 self._nonce = generate_nonce()348 349 headers = [350 (b"Host", request.host.encode("idna")),351 (b"Upgrade", WEBSOCKET_UPGRADE),352 (b"Connection", b"Upgrade"),353 (b"Sec-WebSocket-Key", self._nonce),354 (b"Sec-WebSocket-Version", WEBSOCKET_VERSION),355 ]356 357 if request.subprotocols:358 headers.append(359 (360 b"Sec-WebSocket-Protocol",361 (", ".join(request.subprotocols)).encode("ascii"),362 ),363 )364 365 if request.extensions:366 offers: dict[str, str | bool] = {}367 for e in request.extensions:368 assert isinstance(e, Extension)369 offers[e.name] = e.offer()370 extensions = []371 for name, params in offers.items():372 bname = name.encode("ascii")373 if isinstance(params, bool):374 if params:375 extensions.append(bname)376 else:377 extensions.append(b"%s; %s" % (bname, params.encode("ascii")))378 if extensions:379 headers.append((b"Sec-WebSocket-Extensions", b", ".join(extensions)))380 381 upgrade = h11.Request(382 method=b"GET",383 target=request.target.encode("ascii"),384 headers=headers + request.extra_headers,385 )386 return self._h11_connection.send(upgrade) or b""387 388 def _establish_client_connection(389 self, event: h11.InformationalResponse,390 ) -> AcceptConnection:391 # _establish_client_connection is always called after _initiate_connection.392 assert self._initiating_request is not None393 assert self._nonce is not None394 395 accept = None396 connection_tokens = None397 accepts: list[str] = []398 subprotocol = None399 upgrade = b""400 headers: Headers = []401 for name, value in event.headers:402 name = name.lower()403 if name == b"connection":404 connection_tokens = split_comma_header(value)405 continue # Skip appending to headers406 if name == b"sec-websocket-extensions":407 accepts = split_comma_header(value)408 continue # Skip appending to headers409 if name == b"sec-websocket-accept":410 accept = value411 continue # Skip appending to headers412 if name == b"sec-websocket-protocol":413 subprotocol = value.decode("ascii")414 continue # Skip appending to headers415 if name == b"upgrade":416 upgrade = value417 continue # Skip appending to headers418 headers.append((name, value))419 420 if connection_tokens is None or not any(421 token.lower() == "upgrade" for token in connection_tokens422 ):423 msg = "Missing header, 'Connection: Upgrade'"424 raise RemoteProtocolError(425 msg, event_hint=RejectConnection(),426 )427 if upgrade.lower() != WEBSOCKET_UPGRADE:428 msg = f"Missing header, 'Upgrade: {WEBSOCKET_UPGRADE.decode()}'"429 raise RemoteProtocolError(430 msg,431 event_hint=RejectConnection(),432 )433 accept_token = generate_accept_token(self._nonce)434 if accept != accept_token:435 msg = "Bad accept token"436 raise RemoteProtocolError(msg, event_hint=RejectConnection())437 if subprotocol is not None and subprotocol not in self._initiating_request.subprotocols:438 msg = f"unrecognized subprotocol {subprotocol}"439 raise RemoteProtocolError(440 msg,441 event_hint=RejectConnection(),442 )443 extensions = client_extensions_handshake(444 accepts, cast("Sequence[Extension]", self._initiating_request.extensions),445 )446 447 self._connection = Connection(448 ConnectionType.CLIENT if self.client else ConnectionType.SERVER,449 extensions,450 self._h11_connection.trailing_data[0],451 )452 self._state = ConnectionState.OPEN453 return AcceptConnection(454 extensions=extensions, extra_headers=headers, subprotocol=subprotocol,455 )456 457 def __repr__(self) -> str:458 return f"{self.__class__.__name__}(client={self.client}, state={self.state})"459 460 461def server_extensions_handshake(462 requested: Iterable[str], supported: list[Extension],463) -> bytes | None:464 """465 Agree on the extensions to use returning an appropriate header value.466 467 This returns None if there are no agreed extensions468 """469 accepts: dict[str, bool | bytes] = {}470 for offer in requested:471 name = offer.split(";", 1)[0].strip()472 for extension in supported:473 if extension.name == name:474 accept = extension.accept(offer)475 if isinstance(accept, bool):476 if accept:477 accepts[extension.name] = True478 elif accept is not None:479 accepts[extension.name] = accept.encode("ascii")480 481 if accepts:482 extensions: list[bytes] = []483 for name, params in accepts.items():484 name_bytes = name.encode("ascii")485 if isinstance(params, bool):486 assert params487 extensions.append(name_bytes)488 elif params == b"":489 extensions.append(b"%s" % (name_bytes))490 else:491 extensions.append(b"%s; %s" % (name_bytes, params))492 return b", ".join(extensions)493 494 return None495 496 497def client_extensions_handshake(498 accepted: Iterable[str], supported: Sequence[Extension],499) -> list[Extension]:500 # This raises RemoteProtocolError is the accepted extension is not501 # supported.502 extensions = []503 for accept in accepted:504 name = accept.split(";", 1)[0].strip()505 for extension in supported:506 if extension.name == name:507 extension.finalize(accept)508 extensions.append(extension)509 break510 else:511 msg = f"unrecognized extension {name}"512 raise RemoteProtocolError(513 msg, event_hint=RejectConnection(),514 )515 return extensions516 