Skip to main content

libdd_trace_utils/send_with_retry/
compression.rs

1// Copyright 2024-Present Datadog, Inc. https://www.datadoghq.com/
2// SPDX-License-Identifier: Apache-2.0
3
4#[cfg(feature = "compression")]
5use std::io::Write as _;
6
7#[cfg(feature = "compression")]
8const CONTENT_ENCODING_ZSTD: http::HeaderValue = http::HeaderValue::from_static("zstd");
9#[cfg(feature = "compression")]
10const DEFAULT_COMPRESSION_LEVEL: i32 = 3;
11#[cfg(all(feature = "compression", target_arch = "wasm32"))]
12const MIN_ZRIP_COMPRESSION_LEVEL: i32 = -7;
13#[cfg(all(feature = "compression", target_arch = "wasm32"))]
14const MAX_ZRIP_COMPRESSION_LEVEL: i32 = 4;
15
16#[cfg(all(feature = "compression", not(target_arch = "wasm32")))]
17type ZstdEncoder = zstd::Encoder<'static, Vec<u8>>;
18#[cfg(all(feature = "compression", target_arch = "wasm32"))]
19type ZstdEncoder = zrip::FrameEncoder<Vec<u8>>;
20
21#[cfg(feature = "compression")]
22fn zstd_compression_level(level: i32) -> i32 {
23    let level = if level == 0 {
24        DEFAULT_COMPRESSION_LEVEL
25    } else {
26        level
27    };
28
29    #[cfg(all(feature = "compression", target_arch = "wasm32"))]
30    let level = level.clamp(MIN_ZRIP_COMPRESSION_LEVEL, MAX_ZRIP_COMPRESSION_LEVEL);
31
32    level
33}
34
35#[cfg(all(feature = "compression", not(target_arch = "wasm32")))]
36fn new_zstd_encoder(writer: Vec<u8>, level: i32) -> std::io::Result<ZstdEncoder> {
37    zstd::Encoder::new(writer, level)
38}
39
40#[cfg(all(feature = "compression", target_arch = "wasm32"))]
41fn new_zstd_encoder(writer: Vec<u8>, level: i32) -> std::io::Result<ZstdEncoder> {
42    zrip::FrameEncoder::new(writer, level).map_err(std::io::Error::other)
43}
44
45#[derive(Clone, Copy, Debug)]
46pub enum CompressionStrategy {
47    None,
48    #[cfg(feature = "compression")]
49    /// Zstd-compatible compression.
50    ///
51    /// Native targets accept the range reported by `zstd::compression_level_range()`.
52    /// WASM clamps levels to zrip's supported range of `-7..=4`. Level `0` selects
53    /// level `3` on every target.
54    Zstd {
55        level: i32,
56    },
57}
58
59/// Returns the compressed data, and the actual compression strategy used.
60/// If an error happens during compression, defaults to [`CompressionStrategy::None`]
61pub fn compress(data: Vec<u8>, strategy: CompressionStrategy) -> (Vec<u8>, CompressionStrategy) {
62    match strategy {
63        CompressionStrategy::None => (data, CompressionStrategy::None),
64        #[cfg(feature = "compression")]
65        CompressionStrategy::Zstd { level } => {
66            let level = zstd_compression_level(level);
67            let strategy = CompressionStrategy::Zstd { level };
68            // Start with an initial buffer
69            // Allocate 1/10th of the original buffer, so we shouldn't add too
70            // much memory usage, and no less than 256 bytes
71            let writer = Vec::with_capacity((data.len() / 10).max(256));
72            let result = new_zstd_encoder(writer, level).and_then(|mut encoder| {
73                encoder.write_all(&data)?;
74                Ok((encoder.finish()?, strategy))
75            });
76            result.unwrap_or((data, CompressionStrategy::None))
77        }
78    }
79}
80
81pub fn add_headers(headers: &mut http::HeaderMap, strategy: CompressionStrategy) {
82    match strategy {
83        CompressionStrategy::None => {
84            let _ = headers;
85        }
86        #[cfg(feature = "compression")]
87        CompressionStrategy::Zstd { .. } => {
88            headers.insert(http::header::CONTENT_ENCODING, CONTENT_ENCODING_ZSTD);
89        }
90    }
91}
92
93#[cfg(all(test, feature = "compression", not(target_arch = "wasm32")))]
94mod tests {
95    use super::*;
96
97    fn decompress(data: &[u8]) -> std::io::Result<Vec<u8>> {
98        zstd::decode_all(data)
99    }
100
101    #[test]
102    fn zstd_compression_roundtrips() {
103        let data = b"hello zstd".repeat(100);
104        let (compressed, strategy) = compress(data.clone(), CompressionStrategy::Zstd { level: 1 });
105
106        assert!(matches!(strategy, CompressionStrategy::Zstd { level: 1 }));
107        assert_eq!(decompress(&compressed).unwrap(), data);
108    }
109
110    #[test]
111    fn zero_uses_default_compression_level() {
112        let data = b"hello zstd".repeat(100);
113        let (default_compressed, strategy) =
114            compress(data.clone(), CompressionStrategy::Zstd { level: 0 });
115        let (level_three_compressed, _) = compress(
116            data,
117            CompressionStrategy::Zstd {
118                level: DEFAULT_COMPRESSION_LEVEL,
119            },
120        );
121
122        assert_eq!(default_compressed, level_three_compressed);
123        assert!(matches!(strategy, CompressionStrategy::Zstd { level: 3 }));
124    }
125
126    #[test]
127    fn native_compression_level_is_not_clamped() {
128        let data = b"hello zstd".repeat(100);
129        let (compressed, strategy) =
130            compress(data.clone(), CompressionStrategy::Zstd { level: 22 });
131
132        assert!(matches!(strategy, CompressionStrategy::Zstd { level: 22 }));
133        assert_eq!(decompress(&compressed).unwrap(), data);
134    }
135}