|
1 | 1 | from __future__ import annotations |
2 | 2 |
|
| 3 | +import sys |
3 | 4 | from typing import TYPE_CHECKING, Literal |
4 | 5 |
|
5 | 6 | from litestar.enums import CompressionEncoding |
6 | 7 | from litestar.exceptions import MissingDependencyException |
7 | 8 | from litestar.middleware.compression.facade import CompressionFacade |
8 | 9 |
|
9 | | -try: |
10 | | - import zstandard as zstd |
11 | | -except ImportError as e: |
12 | | - raise MissingDependencyException("zstandard", extra="zstd") from e |
| 10 | +if sys.version_info >= (3, 14): |
| 11 | + from compression import zstd |
| 12 | +else: |
| 13 | + try: |
| 14 | + from backports import zstd |
| 15 | + except ImportError as e: |
| 16 | + raise MissingDependencyException("backports.zstd", extra="zstd") from e |
13 | 17 |
|
14 | 18 | if TYPE_CHECKING: |
15 | 19 | from io import BytesIO |
|
18 | 22 |
|
19 | 23 |
|
20 | 24 | class ZstdCompression(CompressionFacade): |
21 | | - __slots__ = ("buffer", "cctx", "compression_encoding", "compressor") |
| 25 | + __slots__ = ("buffer", "compression_encoding", "compressor") |
22 | 26 |
|
23 | 27 | encoding = CompressionEncoding("zstd") |
24 | | - |
25 | | - def __init__(self, buffer: BytesIO, compression_encoding: Literal["zstd"] | str, config: CompressionConfig) -> None: |
| 28 | + upper_bound = zstd.CompressionParameter.compression_level.bounds()[1] |
| 29 | + |
| 30 | + def __init__( |
| 31 | + self, |
| 32 | + buffer: BytesIO, |
| 33 | + compression_encoding: Literal["zstd"] | str, |
| 34 | + config: CompressionConfig, |
| 35 | + ) -> None: |
26 | 36 | self.buffer = buffer |
27 | 37 | self.compression_encoding = compression_encoding |
28 | | - self.cctx = zstd.ZstdCompressor(level=config.zstd_compress_level) |
29 | | - self.compressor = self.cctx.stream_writer(buffer) |
30 | | - |
31 | | - def write(self, body: bytes | bytearray, final: bool = False) -> None: |
32 | | - self.compressor.write(body) |
33 | | - if final: |
34 | | - self.compressor.flush(zstd.FLUSH_FRAME) |
| 38 | + self.compressor = zstd.ZstdCompressor(level=config.zstd_compress_level) |
| 39 | + |
| 40 | + def write( |
| 41 | + self, |
| 42 | + body: bytes | bytearray, |
| 43 | + final: bool = False, |
| 44 | + ) -> None: |
| 45 | + if not final: |
| 46 | + self.buffer.write(self.compressor.compress(body, mode=zstd.ZstdCompressor.FLUSH_BLOCK)) |
35 | 47 | else: |
36 | | - self.compressor.flush(zstd.FLUSH_BLOCK) |
| 48 | + self.buffer.write(self.compressor.compress(body, mode=zstd.ZstdCompressor.FLUSH_FRAME)) |
37 | 49 |
|
38 | 50 | def close(self) -> None: |
39 | | - self.compressor.flush(zstd.FLUSH_FRAME) |
| 51 | + if self.compressor.last_mode != zstd.ZstdCompressor.FLUSH_FRAME: |
| 52 | + self.buffer.write(self.compressor.flush()) |
0 commit comments