Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
extensions.py321 linesDownload Raw Back to wsproto
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 
codekingpro/portable-devtools · Team Ai