1use postcard::{from_bytes, to_stdvec};
5use serde::{Serialize, de::DeserializeOwned};
6use zstd::{decode_all, encode_all};
7
8use crate::error::{DecodeError, EncodeError};
9
10pub fn encode<T: Serialize + ?Sized>(value: &T, level: i32) -> Result<Vec<u8>, EncodeError> {
11 let raw = to_stdvec(value).map_err(|e| EncodeError::Serialization(e.to_string()))?;
12 encode_all(&raw[..], level).map_err(|e| EncodeError::Compression(e.to_string()))
13}
14
15pub fn decode<T: DeserializeOwned>(bytes: &[u8]) -> Result<T, DecodeError> {
16 let raw = decode_all(bytes).map_err(|e| DecodeError::Decompression(e.to_string()))?;
17 from_bytes(&raw).map_err(|e| DecodeError::Deserialization(e.to_string()))
18}
19
20#[cfg(test)]
21mod tests {
22 use serde::Deserialize;
23
24 use super::*;
25
26 #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
27 struct Sample {
28 version: u64,
29 rows: Vec<String>,
30 }
31
32 fn sample() -> Sample {
33 Sample {
34 version: 42,
35 rows: vec!["alpha".to_string(), "beta".to_string()],
36 }
37 }
38
39 #[test]
40 fn round_trips_a_single_value() {
41 let encoded = encode(&sample(), 1).unwrap();
42 assert_eq!(decode::<Sample>(&encoded).unwrap(), sample());
43 }
44
45 #[test]
46 fn round_trips_a_batch_through_the_same_functions() {
47 let batch = vec![sample(), sample()];
49 let encoded = encode(&batch, 1).unwrap();
50 assert_eq!(decode::<Vec<Sample>>(&encoded).unwrap(), batch);
51 }
52
53 #[test]
54 fn encodes_an_unsized_slice_the_same_as_its_vec() {
55 let batch = vec![sample(), sample()];
57 assert_eq!(encode(&batch[..], 1).unwrap(), encode(&batch, 1).unwrap());
58 }
59
60 #[test]
61 fn every_level_decodes_with_the_same_reader() {
62 for level in [1, 2, 3, 9] {
64 let encoded = encode(&sample(), level).unwrap();
65 assert_eq!(decode::<Sample>(&encoded).unwrap(), sample(), "level {level}");
66 }
67 }
68
69 #[test]
70 fn output_is_compressed_not_raw_postcard() {
71 let big = Sample {
73 version: 1,
74 rows: vec!["x".repeat(4096)],
75 };
76 let raw = to_stdvec(&big).unwrap();
77 assert!(encode(&big, 1).unwrap().len() < raw.len() / 4);
78 }
79
80 #[test]
81 fn rejects_bytes_that_are_not_compressed() {
82 let raw = to_stdvec(&sample()).unwrap();
83 assert!(matches!(decode::<Sample>(&raw), Err(DecodeError::Decompression(_))));
84 }
85
86 #[test]
87 fn rejects_a_payload_that_decompresses_to_the_wrong_type() {
88 let encoded = encode(&"not a sample".to_string(), 1).unwrap();
89 assert!(matches!(decode::<Sample>(&encoded), Err(DecodeError::Deserialization(_))));
90 }
91
92 #[test]
93 fn rejects_truncated_input() {
94 let encoded = encode(&sample(), 1).unwrap();
95 let truncated = &encoded[..encoded.len() / 2];
96 assert!(decode::<Sample>(truncated).is_err());
97 }
98}