Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
permessage_deflate.py700 linesDownload Raw Back to extensions
1from __future__ import annotations
2
3import zlib
4from collections.abc import Sequence
5from typing import Any, Literal
6
7from .. import frames
8from ..exceptions import (
9    DuplicateParameter,
10    InvalidParameterName,
11    InvalidParameterValue,
12    NegotiationError,
13    PayloadTooBig,
14    ProtocolError,
15)
16from ..typing import BytesLike, ExtensionName, ExtensionParameter
17from .base import ClientExtensionFactory, Extension, ServerExtensionFactory
18
19
20__all__ = [
21    "PerMessageDeflate",
22    "ClientPerMessageDeflateFactory",
23    "enable_client_permessage_deflate",
24    "ServerPerMessageDeflateFactory",
25    "enable_server_permessage_deflate",
26]
27
28_EMPTY_UNCOMPRESSED_BLOCK = b"\x00\x00\xff\xff"
29
30_MAX_WINDOW_BITS_VALUES = [str(bits) for bits in range(8, 16)]
31
32
33class PerMessageDeflate(Extension):
34    """
35    Per-Message Deflate extension.
36
37    """
38
39    name = ExtensionName("permessage-deflate")
40
41    def __init__(
42        self,
43        remote_no_context_takeover: bool,
44        local_no_context_takeover: bool,
45        remote_max_window_bits: int,
46        local_max_window_bits: int,
47        compress_settings: dict[Any, Any] | None = None,
48    ) -> None:
49        """
50        Configure the Per-Message Deflate extension.
51
52        """
53        if compress_settings is None:
54            compress_settings = {}
55
56        assert remote_no_context_takeover in [False, True]
57        assert local_no_context_takeover in [False, True]
58        assert 8 <= remote_max_window_bits <= 15
59        assert 8 <= local_max_window_bits <= 15
60        assert "wbits" not in compress_settings
61
62        self.remote_no_context_takeover = remote_no_context_takeover
63        self.local_no_context_takeover = local_no_context_takeover
64        self.remote_max_window_bits = remote_max_window_bits
65        self.local_max_window_bits = local_max_window_bits
66        self.compress_settings = compress_settings
67
68        if not self.remote_no_context_takeover:
69            self.decoder = zlib.decompressobj(wbits=-self.remote_max_window_bits)
70
71        if not self.local_no_context_takeover:
72            self.encoder = zlib.compressobj(
73                wbits=-self.local_max_window_bits,
74                **self.compress_settings,
75            )
76
77        # To handle continuation frames properly, we must keep track of
78        # whether that initial frame was encoded.
79        self.decode_cont_data = False
80        # There's no need for self.encode_cont_data because we always encode
81        # outgoing frames, so it would always be True.
82
83    def __repr__(self) -> str:
84        return (
85            f"PerMessageDeflate("
86            f"remote_no_context_takeover={self.remote_no_context_takeover}, "
87            f"local_no_context_takeover={self.local_no_context_takeover}, "
88            f"remote_max_window_bits={self.remote_max_window_bits}, "
89            f"local_max_window_bits={self.local_max_window_bits})"
90        )
91
92    def decode(
93        self,
94        frame: frames.Frame,
95        *,
96        max_size: int | None = None,
97    ) -> frames.Frame:
98        """
99        Decode an incoming frame.
100
101        """
102        # Skip control frames.
103        if frame.opcode in frames.CTRL_OPCODES:
104            return frame
105
106        # Handle continuation data frames:
107        # - skip if the message isn't encoded
108        # - reset "decode continuation data" flag if it's a final frame
109        if frame.opcode is frames.OP_CONT:
110            if not self.decode_cont_data:
111                return frame
112            if frame.fin:
113                self.decode_cont_data = False
114
115        # Handle text and binary data frames:
116        # - skip if the message isn't encoded
117        # - unset the rsv1 flag on the first frame of a compressed message
118        # - set "decode continuation data" flag if it's a non-final frame
119        else:
120            if not frame.rsv1:
121                return frame
122            if not frame.fin:
123                self.decode_cont_data = True
124
125            # Re-initialize per-message decoder.
126            if self.remote_no_context_takeover:
127                self.decoder = zlib.decompressobj(wbits=-self.remote_max_window_bits)
128
129        # Uncompress data. Protect against zip bombs by preventing zlib from
130        # decompressing more than max_length bytes (except when the limit is
131        # disabled with max_size = None).
132        data: BytesLike
133        if frame.fin and len(frame.data) < 2044:
134            # Profiling shows that appending four bytes, which makes a copy, is
135            # faster than calling decompress() again when data is less than 2kB.
136            data = bytes(frame.data) + _EMPTY_UNCOMPRESSED_BLOCK
137        else:
138            data = frame.data
139        max_length = 0 if max_size is None else max_size
140        try:
141            data = self.decoder.decompress(data, max_length)
142            if self.decoder.unconsumed_tail:
143                assert max_size is not None  # help mypy
144                raise PayloadTooBig(None, max_size)
145            if frame.fin and len(frame.data) >= 2044:
146                # This cannot generate additional data.
147                self.decoder.decompress(_EMPTY_UNCOMPRESSED_BLOCK)
148        except zlib.error as exc:
149            raise ProtocolError("decompression failed") from exc
150
151        # Allow garbage collection of the decoder if it won't be reused.
152        if frame.fin and self.remote_no_context_takeover:
153            del self.decoder
154
155        return frames.Frame(
156            frame.opcode,
157            data,
158            frame.fin,
159            # Unset the rsv1 flag on the first frame of a compressed message.
160            False,
161            frame.rsv2,
162            frame.rsv3,
163        )
164
165    def encode(self, frame: frames.Frame) -> frames.Frame:
166        """
167        Encode an outgoing frame.
168
169        """
170        # Skip control frames.
171        if frame.opcode in frames.CTRL_OPCODES:
172            return frame
173
174        # Since we always encode messages, there's no "encode continuation
175        # data" flag similar to "decode continuation data" at this time.
176
177        if frame.opcode is not frames.OP_CONT:
178            # Re-initialize per-message decoder.
179            if self.local_no_context_takeover:
180                self.encoder = zlib.compressobj(
181                    wbits=-self.local_max_window_bits,
182                    **self.compress_settings,
183                )
184
185        # Compress data.
186        data: BytesLike
187        data = self.encoder.compress(frame.data) + self.encoder.flush(zlib.Z_SYNC_FLUSH)
188        if frame.fin:
189            # Sync flush generates between 5 or 6 bytes, ending with the bytes
190            # 0x00 0x00 0xff 0xff, which must be removed.
191            assert data[-4:] == _EMPTY_UNCOMPRESSED_BLOCK
192            # Making a copy is faster than memoryview(a)[:-4] until 2kB.
193            if len(data) < 2048:
194                data = data[:-4]
195            else:
196                data = memoryview(data)[:-4]
197
198        # Allow garbage collection of the encoder if it won't be reused.
199        if frame.fin and self.local_no_context_takeover:
200            del self.encoder
201
202        return frames.Frame(
203            frame.opcode,
204            data,
205            frame.fin,
206            # Set the rsv1 flag on the first frame of a compressed message.
207            frame.opcode is not frames.OP_CONT,
208            frame.rsv2,
209            frame.rsv3,
210        )
211
212
213def _build_parameters(
214    server_no_context_takeover: bool,
215    client_no_context_takeover: bool,
216    server_max_window_bits: int | None,
217    client_max_window_bits: int | Literal[True] | None,
218) -> list[ExtensionParameter]:
219    """
220    Build a list of ``(name, value)`` pairs for some compression parameters.
221
222    """
223    params: list[ExtensionParameter] = []
224    if server_no_context_takeover:
225        params.append(("server_no_context_takeover", None))
226    if client_no_context_takeover:
227        params.append(("client_no_context_takeover", None))
228    if server_max_window_bits:
229        params.append(("server_max_window_bits", str(server_max_window_bits)))
230    if client_max_window_bits is True:  # only in handshake requests
231        params.append(("client_max_window_bits", None))
232    elif client_max_window_bits:
233        params.append(("client_max_window_bits", str(client_max_window_bits)))
234    return params
235
236
237def _extract_parameters(
238    params: Sequence[ExtensionParameter], *, is_server: bool
239) -> tuple[bool, bool, int | None, int | Literal[True] | None]:
240    """
241    Extract compression parameters from a list of ``(name, value)`` pairs.
242
243    If ``is_server`` is :obj:`True`, ``client_max_window_bits`` may be
244    provided without a value. This is only allowed in handshake requests.
245
246    """
247    server_no_context_takeover: bool = False
248    client_no_context_takeover: bool = False
249    server_max_window_bits: int | None = None
250    client_max_window_bits: int | Literal[True] | None = None
251
252    for name, value in params:
253        if name == "server_no_context_takeover":
254            if server_no_context_takeover:
255                raise DuplicateParameter(name)
256            if value is None:
257                server_no_context_takeover = True
258            else:
259                raise InvalidParameterValue(name, value)
260
261        elif name == "client_no_context_takeover":
262            if client_no_context_takeover:
263                raise DuplicateParameter(name)
264            if value is None:
265                client_no_context_takeover = True
266            else:
267                raise InvalidParameterValue(name, value)
268
269        elif name == "server_max_window_bits":
270            if server_max_window_bits is not None:
271                raise DuplicateParameter(name)
272            if value in _MAX_WINDOW_BITS_VALUES:
273                server_max_window_bits = int(value)
274            else:
275                raise InvalidParameterValue(name, value)
276
277        elif name == "client_max_window_bits":
278            if client_max_window_bits is not None:
279                raise DuplicateParameter(name)
280            if is_server and value is None:  # only in handshake requests
281                client_max_window_bits = True
282            elif value in _MAX_WINDOW_BITS_VALUES:
283                client_max_window_bits = int(value)
284            else:
285                raise InvalidParameterValue(name, value)
286
287        else:
288            raise InvalidParameterName(name)
289
290    return (
291        server_no_context_takeover,
292        client_no_context_takeover,
293        server_max_window_bits,
294        client_max_window_bits,
295    )
296
297
298class ClientPerMessageDeflateFactory(ClientExtensionFactory):
299    """
300    Client-side extension factory for the Per-Message Deflate extension.
301
302    Parameters behave as described in `section 7.1 of RFC 7692`_.
303
304    .. _section 7.1 of RFC 7692: https://datatracker.ietf.org/doc/html/rfc7692#section-7.1
305
306    Set them to :obj:`True` to include them in the negotiation offer without a
307    value or to an integer value to include them with this value.
308
309    Args:
310        server_no_context_takeover: Prevent server from using context takeover.
311        client_no_context_takeover: Prevent client from using context takeover.
312        server_max_window_bits: Maximum size of the server's LZ77 sliding window
313            in bits, between 8 and 15.
314        client_max_window_bits: Maximum size of the client's LZ77 sliding window
315            in bits, between 8 and 15, or :obj:`True` to indicate support without
316            setting a limit.
317        compress_settings: Additional keyword arguments for :func:`zlib.compressobj`,
318            excluding ``wbits``.
319
320    """
321
322    name = ExtensionName("permessage-deflate")
323
324    def __init__(
325        self,
326        server_no_context_takeover: bool = False,
327        client_no_context_takeover: bool = False,
328        server_max_window_bits: int | None = None,
329        client_max_window_bits: int | Literal[True] | None = True,
330        compress_settings: dict[str, Any] | None = None,
331    ) -> None:
332        """
333        Configure the Per-Message Deflate extension factory.
334
335        """
336        if not (server_max_window_bits is None or 8 <= server_max_window_bits <= 15):
337            raise ValueError("server_max_window_bits must be between 8 and 15")
338        if not (
339            client_max_window_bits is None
340            or client_max_window_bits is True
341            or 8 <= client_max_window_bits <= 15
342        ):
343            raise ValueError("client_max_window_bits must be between 8 and 15")
344        if compress_settings is not None and "wbits" in compress_settings:
345            raise ValueError(
346                "compress_settings must not include wbits, "
347                "set client_max_window_bits instead"
348            )
349
350        self.server_no_context_takeover = server_no_context_takeover
351        self.client_no_context_takeover = client_no_context_takeover
352        self.server_max_window_bits = server_max_window_bits
353        self.client_max_window_bits = client_max_window_bits
354        self.compress_settings = compress_settings
355
356    def get_request_params(self) -> Sequence[ExtensionParameter]:
357        """
358        Build request parameters.
359
360        """
361        return _build_parameters(
362            self.server_no_context_takeover,
363            self.client_no_context_takeover,
364            self.server_max_window_bits,
365            self.client_max_window_bits,
366        )
367
368    def process_response_params(
369        self,
370        params: Sequence[ExtensionParameter],
371        accepted_extensions: Sequence[Extension],
372    ) -> PerMessageDeflate:
373        """
374        Process response parameters.
375
376        Return an extension instance.
377
378        """
379        if any(other.name == self.name for other in accepted_extensions):
380            raise NegotiationError(f"received duplicate {self.name}")
381
382        # Request parameters are available in instance variables.
383
384        # Load response parameters in local variables.
385        (
386            server_no_context_takeover,
387            client_no_context_takeover,
388            server_max_window_bits,
389            client_max_window_bits,
390        ) = _extract_parameters(params, is_server=False)
391
392        # After comparing the request and the response, the final
393        # configuration must be available in the local variables.
394
395        # server_no_context_takeover
396        #
397        #   Req.    Resp.   Result
398        #   ------  ------  --------------------------------------------------
399        #   False   False   False
400        #   False   True    True
401        #   True    False   Error!
402        #   True    True    True
403
404        if self.server_no_context_takeover:
405            if not server_no_context_takeover:
406                raise NegotiationError("expected server_no_context_takeover")
407
408        # client_no_context_takeover
409        #
410        #   Req.    Resp.   Result
411        #   ------  ------  --------------------------------------------------
412        #   False   False   False
413        #   False   True    True
414        #   True    False   True - must change value
415        #   True    True    True
416
417        if self.client_no_context_takeover:
418            if not client_no_context_takeover:
419                client_no_context_takeover = True
420
421        # server_max_window_bits
422
423        #   Req.    Resp.   Result
424        #   ------  ------  --------------------------------------------------
425        #   None    None    None
426        #   None    8≤M≤15  M
427        #   8≤N≤15  None    Error!
428        #   8≤N≤15  8≤M≤N   M
429        #   8≤N≤15  N<M≤15  Error!
430
431        if self.server_max_window_bits is None:
432            pass
433
434        else:
435            if server_max_window_bits is None:
436                raise NegotiationError("expected server_max_window_bits")
437            elif server_max_window_bits > self.server_max_window_bits:
438                raise NegotiationError("unsupported server_max_window_bits")
439
440        # client_max_window_bits
441
442        #   Req.    Resp.   Result
443        #   ------  ------  --------------------------------------------------
444        #   None    None    None
445        #   None    8≤M≤15  Error!
446        #   True    None    None
447        #   True    8≤M≤15  M
448        #   8≤N≤15  None    N - must change value
449        #   8≤N≤15  8≤M≤N   M
450        #   8≤N≤15  N<M≤15  Error!
451
452        if self.client_max_window_bits is None:
453            if client_max_window_bits is not None:
454                raise NegotiationError("unexpected client_max_window_bits")
455
456        elif self.client_max_window_bits is True:
457            pass
458
459        else:
460            if client_max_window_bits is None:
461                client_max_window_bits = self.client_max_window_bits
462            elif client_max_window_bits > self.client_max_window_bits:
463                raise NegotiationError("unsupported client_max_window_bits")
464
465        return PerMessageDeflate(
466            server_no_context_takeover,  # remote_no_context_takeover
467            client_no_context_takeover,  # local_no_context_takeover
468            server_max_window_bits or 15,  # remote_max_window_bits
469            client_max_window_bits or 15,  # local_max_window_bits
470            self.compress_settings,
471        )
472
473
474def enable_client_permessage_deflate(
475    extensions: Sequence[ClientExtensionFactory] | None,
476) -> Sequence[ClientExtensionFactory]:
477    """
478    Enable Per-Message Deflate with default settings in client extensions.
479
480    If the extension is already present, perhaps with non-default settings,
481    the configuration isn't changed.
482
483    """
484    if extensions is None:
485        extensions = []
486    if not any(
487        extension_factory.name == ClientPerMessageDeflateFactory.name
488        for extension_factory in extensions
489    ):
490        extensions = list(extensions) + [
491            ClientPerMessageDeflateFactory(
492                compress_settings={"memLevel": 5},
493            )
494        ]
495    return extensions
496
497
498class ServerPerMessageDeflateFactory(ServerExtensionFactory):
499    """
500    Server-side extension factory for the Per-Message Deflate extension.
501
502    Parameters behave as described in `section 7.1 of RFC 7692`_.
503
504    .. _section 7.1 of RFC 7692: https://datatracker.ietf.org/doc/html/rfc7692#section-7.1
505
506    Set them to :obj:`True` to include them in the negotiation offer without a
507    value or to an integer value to include them with this value.
508
509    Args:
510        server_no_context_takeover: Prevent server from using context takeover.
511        client_no_context_takeover: Prevent client from using context takeover.
512        server_max_window_bits: Maximum size of the server's LZ77 sliding window
513            in bits, between 8 and 15.
514        client_max_window_bits: Maximum size of the client's LZ77 sliding window
515            in bits, between 8 and 15.
516        compress_settings: Additional keyword arguments for :func:`zlib.compressobj`,
517            excluding ``wbits``.
518        require_client_max_window_bits: Do not enable compression at all if
519            client doesn't advertise support for ``client_max_window_bits``;
520            the default behavior is to enable compression without enforcing
521            ``client_max_window_bits``.
522
523    """
524
525    name = ExtensionName("permessage-deflate")
526
527    def __init__(
528        self,
529        server_no_context_takeover: bool = False,
530        client_no_context_takeover: bool = False,
531        server_max_window_bits: int | None = None,
532        client_max_window_bits: int | None = None,
533        compress_settings: dict[str, Any] | None = None,
534        require_client_max_window_bits: bool = False,
535    ) -> None:
536        """
537        Configure the Per-Message Deflate extension factory.
538
539        """
540        if not (server_max_window_bits is None or 8 <= server_max_window_bits <= 15):
541            raise ValueError("server_max_window_bits must be between 8 and 15")
542        if not (client_max_window_bits is None or 8 <= client_max_window_bits <= 15):
543            raise ValueError("client_max_window_bits must be between 8 and 15")
544        if compress_settings is not None and "wbits" in compress_settings:
545            raise ValueError(
546                "compress_settings must not include wbits, "
547                "set server_max_window_bits instead"
548            )
549        if client_max_window_bits is None and require_client_max_window_bits:
550            raise ValueError(
551                "require_client_max_window_bits is enabled, "
552                "but client_max_window_bits isn't configured"
553            )
554
555        self.server_no_context_takeover = server_no_context_takeover
556        self.client_no_context_takeover = client_no_context_takeover
557        self.server_max_window_bits = server_max_window_bits
558        self.client_max_window_bits = client_max_window_bits
559        self.compress_settings = compress_settings
560        self.require_client_max_window_bits = require_client_max_window_bits
561
562    def process_request_params(
563        self,
564        params: Sequence[ExtensionParameter],
565        accepted_extensions: Sequence[Extension],
566    ) -> tuple[list[ExtensionParameter], PerMessageDeflate]:
567        """
568        Process request parameters.
569
570        Return response params and an extension instance.
571
572        """
573        if any(other.name == self.name for other in accepted_extensions):
574            raise NegotiationError(f"skipped duplicate {self.name}")
575
576        # Load request parameters in local variables.
577        (
578            server_no_context_takeover,
579            client_no_context_takeover,
580            server_max_window_bits,
581            client_max_window_bits,
582        ) = _extract_parameters(params, is_server=True)
583
584        # Configuration parameters are available in instance variables.
585
586        # After comparing the request and the configuration, the response must
587        # be available in the local variables.
588
589        # server_no_context_takeover
590        #
591        #   Config  Req.    Resp.
592        #   ------  ------  --------------------------------------------------
593        #   False   False   False
594        #   False   True    True
595        #   True    False   True - must change value to True
596        #   True    True    True
597
598        if self.server_no_context_takeover:
599            if not server_no_context_takeover:
600                server_no_context_takeover = True
601
602        # client_no_context_takeover
603        #
604        #   Config  Req.    Resp.
605        #   ------  ------  --------------------------------------------------
606        #   False   False   False
607        #   False   True    True (or False)
608        #   True    False   True - must change value to True
609        #   True    True    True (or False)
610
611        if self.client_no_context_takeover:
612            if not client_no_context_takeover:
613                client_no_context_takeover = True
614
615        # server_max_window_bits
616
617        #   Config  Req.    Resp.
618        #   ------  ------  --------------------------------------------------
619        #   None    None    None
620        #   None    8≤M≤15  M
621        #   8≤N≤15  None    N - must change value
622        #   8≤N≤15  8≤M≤N   M
623        #   8≤N≤15  N<M≤15  N - must change value
624
625        if self.server_max_window_bits is None:
626            pass
627
628        else:
629            if server_max_window_bits is None:
630                server_max_window_bits = self.server_max_window_bits
631            elif server_max_window_bits > self.server_max_window_bits:
632                server_max_window_bits = self.server_max_window_bits
633
634        # client_max_window_bits
635
636        #   Config  Req.    Resp.
637        #   ------  ------  --------------------------------------------------
638        #   None    None    None
639        #   None    True    None - must change value
640        #   None    8≤M≤15  M (or None)
641        #   8≤N≤15  None    None or Error!
642        #   8≤N≤15  True    N - must change value
643        #   8≤N≤15  8≤M≤N   M (or None)
644        #   8≤N≤15  N<M≤15  N
645
646        if self.client_max_window_bits is None:
647            if client_max_window_bits is True:
648                client_max_window_bits = self.client_max_window_bits
649
650        else:
651            if client_max_window_bits is None:
652                if self.require_client_max_window_bits:
653                    raise NegotiationError("required client_max_window_bits")
654            elif client_max_window_bits is True:
655                client_max_window_bits = self.client_max_window_bits
656            elif self.client_max_window_bits < client_max_window_bits:
657                client_max_window_bits = self.client_max_window_bits
658
659        return (
660            _build_parameters(
661                server_no_context_takeover,
662                client_no_context_takeover,
663                server_max_window_bits,
664                client_max_window_bits,
665            ),
666            PerMessageDeflate(
667                client_no_context_takeover,  # remote_no_context_takeover
668                server_no_context_takeover,  # local_no_context_takeover
669                client_max_window_bits or 15,  # remote_max_window_bits
670                server_max_window_bits or 15,  # local_max_window_bits
671                self.compress_settings,
672            ),
673        )
674
675
676def enable_server_permessage_deflate(
677    extensions: Sequence[ServerExtensionFactory] | None,
678) -> Sequence[ServerExtensionFactory]:
679    """
680    Enable Per-Message Deflate with default settings in server extensions.
681
682    If the extension is already present, perhaps with non-default settings,
683    the configuration isn't changed.
684
685    """
686    if extensions is None:
687        extensions = []
688    if not any(
689        ext_factory.name == ServerPerMessageDeflateFactory.name
690        for ext_factory in extensions
691    ):
692        extensions = list(extensions) + [
693            ServerPerMessageDeflateFactory(
694                server_max_window_bits=12,
695                client_max_window_bits=12,
696                compress_settings={"memLevel": 5},
697            )
698        ]
699    return extensions
700 
codekingpro/portable-devtools · Team Ai