codekingpro/portable-devtools
114k
1"""2wsproto/extensions3~~~~~~~~~~~~~~~~~~4 5WebSocket extensions.6"""7from __future__ import annotations8 9import zlib10from abc import ABC, abstractmethod11from typing import Optional12 13from .frame_protocol import CloseReason, FrameDecoder, FrameProtocol, Opcode, RsvBits14 15 16class Extension(ABC):17 name: str18 19 def enabled(self) -> bool:20 return False21 22 @abstractmethod23 def offer(self) -> bool | str:24 pass25 26 def accept(self, offer: str) -> bool | str | None:27 pass28 29 def finalize(self, offer: str) -> None:30 pass31 32 def frame_inbound_header(33 self,34 proto: FrameDecoder | FrameProtocol,35 opcode: Opcode,36 rsv: RsvBits,37 payload_length: int,38 ) -> CloseReason | RsvBits:39 return RsvBits(False, False, False)40 41 def frame_inbound_payload_data(42 self, proto: FrameDecoder | FrameProtocol, data: bytes,43 ) -> bytes | CloseReason:44 return data45 46 def frame_inbound_complete(47 self, proto: FrameDecoder | FrameProtocol, fin: bool,48 ) -> bytes | CloseReason | None:49 pass50 51 def frame_outbound(52 self,53 proto: FrameDecoder | FrameProtocol,54 opcode: Opcode,55 rsv: RsvBits,56 data: bytes,57 fin: bool,58 ) -> tuple[RsvBits, bytes]:59 return (rsv, data)60 61 62class PerMessageDeflate(Extension):63 name = "permessage-deflate"64 65 DEFAULT_CLIENT_MAX_WINDOW_BITS = 1566 DEFAULT_SERVER_MAX_WINDOW_BITS = 1567 68 def __init__(69 self,70 client_no_context_takeover: bool = False,71 client_max_window_bits: int | None = None,72 server_no_context_takeover: bool = False,73 server_max_window_bits: int | None = None,74 ) -> None:75 self.client_no_context_takeover = client_no_context_takeover76 self.server_no_context_takeover = server_no_context_takeover77 self._client_max_window_bits = self.DEFAULT_CLIENT_MAX_WINDOW_BITS78 self._server_max_window_bits = self.DEFAULT_SERVER_MAX_WINDOW_BITS79 if client_max_window_bits is not None:80 self.client_max_window_bits = client_max_window_bits81 if server_max_window_bits is not None:82 self.server_max_window_bits = server_max_window_bits83 84 self._compressor: Optional[zlib._Compress] = None # noqa85 self._decompressor: Optional[zlib._Decompress] = None # noqa86 # This refers to the current frame87 self._inbound_is_compressible: bool | None = None88 # This refers to the ongoing message (which might span multiple89 # frames). Only the first frame in a fragmented message is flagged for90 # compression, so this carries that bit forward.91 self._inbound_compressed: bool | None = None92 93 self._enabled = False94 95 @property96 def client_max_window_bits(self) -> int:97 return self._client_max_window_bits98 99 @client_max_window_bits.setter100 def client_max_window_bits(self, value: int) -> None:101 if value < 9 or value > 15:102 msg = "Window size must be between 9 and 15 inclusive"103 raise ValueError(msg)104 self._client_max_window_bits = value105 106 @property107 def server_max_window_bits(self) -> int:108 return self._server_max_window_bits109 110 @server_max_window_bits.setter111 def server_max_window_bits(self, value: int) -> None:112 if value < 9 or value > 15:113 msg = "Window size must be between 9 and 15 inclusive"114 raise ValueError(msg)115 self._server_max_window_bits = value116 117 def _compressible_opcode(self, opcode: Opcode) -> bool:118 return opcode in (Opcode.TEXT, Opcode.BINARY, Opcode.CONTINUATION)119 120 def enabled(self) -> bool:121 return self._enabled122 123 def offer(self) -> bool | str:124 parameters = [125 f"client_max_window_bits={self.client_max_window_bits}",126 f"server_max_window_bits={self.server_max_window_bits}",127 ]128 129 if self.client_no_context_takeover:130 parameters.append("client_no_context_takeover")131 if self.server_no_context_takeover:132 parameters.append("server_no_context_takeover")133 134 return "; ".join(parameters)135 136 def finalize(self, offer: str) -> None:137 bits = [b.strip() for b in offer.split(";")]138 for bit in bits[1:]:139 if bit.startswith("client_no_context_takeover"):140 self.client_no_context_takeover = True141 elif bit.startswith("server_no_context_takeover"):142 self.server_no_context_takeover = True143 elif bit.startswith("client_max_window_bits"):144 self.client_max_window_bits = int(bit.split("=", 1)[1].strip())145 elif bit.startswith("server_max_window_bits"):146 self.server_max_window_bits = int(bit.split("=", 1)[1].strip())147 148 self._enabled = True149 150 def _parse_params(self, params: str) -> tuple[int | None, int | None]:151 client_max_window_bits = None152 server_max_window_bits = None153 154 bits = [b.strip() for b in params.split(";")]155 for bit in bits[1:]:156 if bit.startswith("client_no_context_takeover"):157 self.client_no_context_takeover = True158 elif bit.startswith("server_no_context_takeover"):159 self.server_no_context_takeover = True160 elif bit.startswith("client_max_window_bits"):161 if "=" in bit:162 client_max_window_bits = int(bit.split("=", 1)[1].strip())163 else:164 client_max_window_bits = self.client_max_window_bits165 elif bit.startswith("server_max_window_bits"):166 if "=" in bit:167 server_max_window_bits = int(bit.split("=", 1)[1].strip())168 else:169 server_max_window_bits = self.server_max_window_bits170 171 return client_max_window_bits, server_max_window_bits172 173 def accept(self, offer: str) -> bool | None | str:174 client_max_window_bits, server_max_window_bits = self._parse_params(offer)175 176 parameters = []177 178 if self.client_no_context_takeover:179 parameters.append("client_no_context_takeover")180 if self.server_no_context_takeover:181 parameters.append("server_no_context_takeover")182 try:183 if client_max_window_bits is not None:184 parameters.append(f"client_max_window_bits={client_max_window_bits}")185 self.client_max_window_bits = client_max_window_bits186 if server_max_window_bits is not None:187 parameters.append(f"server_max_window_bits={server_max_window_bits}")188 self.server_max_window_bits = server_max_window_bits189 except ValueError:190 return None191 else:192 self._enabled = True193 return "; ".join(parameters)194 195 def frame_inbound_header(196 self,197 proto: FrameDecoder | FrameProtocol,198 opcode: Opcode,199 rsv: RsvBits,200 payload_length: int,201 ) -> CloseReason | RsvBits:202 if rsv.rsv1 and opcode.iscontrol():203 return CloseReason.PROTOCOL_ERROR204 if rsv.rsv1 and opcode is Opcode.CONTINUATION:205 return CloseReason.PROTOCOL_ERROR206 207 self._inbound_is_compressible = self._compressible_opcode(opcode)208 209 if self._inbound_compressed is None:210 self._inbound_compressed = rsv.rsv1211 if self._inbound_compressed:212 assert self._inbound_is_compressible213 if proto.client:214 bits = self.server_max_window_bits215 else:216 bits = self.client_max_window_bits217 if self._decompressor is None:218 self._decompressor = zlib.decompressobj(-int(bits))219 220 return RsvBits(True, False, False)221 222 def frame_inbound_payload_data(223 self, proto: FrameDecoder | FrameProtocol, data: bytes,224 ) -> bytes | CloseReason:225 if not self._inbound_compressed or not self._inbound_is_compressible:226 return data227 assert self._decompressor is not None228 229 try:230 return self._decompressor.decompress(bytes(data))231 except zlib.error:232 return CloseReason.INVALID_FRAME_PAYLOAD_DATA233 234 def frame_inbound_complete(235 self, proto: FrameDecoder | FrameProtocol, fin: bool,236 ) -> bytes | CloseReason | None:237 if not fin:238 return None239 if not self._inbound_is_compressible:240 self._inbound_compressed = None241 return None242 if not self._inbound_compressed:243 self._inbound_compressed = None244 return None245 assert self._decompressor is not None246 247 try:248 data = self._decompressor.decompress(b"\x00\x00\xff\xff")249 data += self._decompressor.flush()250 except zlib.error:251 return CloseReason.INVALID_FRAME_PAYLOAD_DATA252 253 if proto.client:254 no_context_takeover = self.server_no_context_takeover255 else:256 no_context_takeover = self.client_no_context_takeover257 258 if no_context_takeover:259 self._decompressor = None260 261 self._inbound_compressed = None262 263 return data264 265 def frame_outbound(266 self,267 proto: FrameDecoder | FrameProtocol,268 opcode: Opcode,269 rsv: RsvBits,270 data: bytes,271 fin: bool,272 ) -> tuple[RsvBits, bytes]:273 if not self._compressible_opcode(opcode):274 return (rsv, data)275 276 if opcode is not Opcode.CONTINUATION:277 rsv = RsvBits(True, rsv[1], rsv[2])278 279 if self._compressor is None:280 assert opcode is not Opcode.CONTINUATION281 if proto.client:282 bits = self.client_max_window_bits283 else:284 bits = self.server_max_window_bits285 self._compressor = zlib.compressobj(286 zlib.Z_DEFAULT_COMPRESSION, zlib.DEFLATED, -int(bits),287 )288 289 data = self._compressor.compress(bytes(data))290 291 if fin:292 data += self._compressor.flush(zlib.Z_SYNC_FLUSH)293 data = data[:-4]294 295 if proto.client:296 no_context_takeover = self.client_no_context_takeover297 else:298 no_context_takeover = self.server_no_context_takeover299 300 if no_context_takeover:301 self._compressor = None302 303 return (rsv, data)304 305 def __repr__(self) -> str:306 descr = [f"client_max_window_bits={self.client_max_window_bits}"]307 if self.client_no_context_takeover:308 descr.append("client_no_context_takeover")309 descr.append(f"server_max_window_bits={self.server_max_window_bits}")310 if self.server_no_context_takeover:311 descr.append("server_no_context_takeover")312 313 return "<{} {}>".format(self.__class__.__name__, "; ".join(descr))314 315 316#: SUPPORTED_EXTENSIONS maps all supported extension names to their class.317#: This can be used to iterate all supported extensions of wsproto, instantiate318#: new extensions based on their name, or check if a given extension is319#: supported or not.320SUPPORTED_EXTENSIONS = {PerMessageDeflate.name: PerMessageDeflate}321 