codekingpro/portable-devtools
114k
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 