1#![warn(missing_docs)]
4
5use std::io::{Read, Write};
6
7use data_encoding::BASE64;
8use serde::{Deserialize, Serialize};
9
10#[derive(Debug, Clone, PartialEq, Eq, Default, Hash)]
12pub struct Buffer(Vec<u8>);
13
14impl Buffer {
15 pub fn new(data: Vec<u8>) -> Self {
17 Self(data)
18 }
19
20 pub fn into_vec(self) -> Vec<u8> {
22 self.0
23 }
24}
25
26impl<'de> Deserialize<'de> for Buffer {
27 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
28 where
29 D: serde::Deserializer<'de>,
30 {
31 #[derive(Deserialize)]
32 enum BufferData {
33 #[serde(rename = "base64")]
34 Base64(String),
35 #[serde(rename = "zbase64")]
36 ZBase64(String),
37 }
38
39 #[derive(Deserialize)]
40 struct BufferInner {
41 t: String,
42 #[serde(flatten)]
43 data: BufferData,
44 }
45
46 let BufferInner { t, data } = BufferInner::deserialize(deserializer)?;
47
48 if t != "buffer" {
49 return Err(serde::de::Error::custom("expected buffer"));
50 }
51
52 let data = match data {
53 BufferData::Base64(base64) => BASE64
54 .decode(base64.as_bytes())
55 .map_err(serde::de::Error::custom)?,
56 BufferData::ZBase64(zbase64) => {
57 let compressed = BASE64
58 .decode(zbase64.as_bytes())
59 .map_err(serde::de::Error::custom)?;
60 let mut decoder = zstd::stream::Decoder::new(&compressed[..])
61 .map_err(serde::de::Error::custom)?;
62 let mut data = Vec::new();
63 decoder
64 .read_to_end(&mut data)
65 .map_err(serde::de::Error::custom)?;
66 data
67 }
68 };
69
70 Ok(Self(data))
71 }
72}
73
74impl Serialize for Buffer {
75 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
76 where
77 S: serde::ser::Serializer,
78 {
79 use serde::ser::SerializeMap;
80
81 let mut map = serializer.serialize_map(Some(3))?;
82
83 map.serialize_entry("m", &())?; map.serialize_entry("t", "buffer")?;
85
86 if false {
89 let base64 = BASE64.encode(&self.0);
90
91 let mut compressed: Vec<u8> = Vec::new();
92 let mut encoder = zstd::stream::Encoder::new(&mut compressed, 0).unwrap();
93 encoder
94 .set_pledged_src_size(Some(self.0.len() as u64))
95 .unwrap();
96 encoder.include_contentsize(true).unwrap();
97 encoder.write_all(&self.0).unwrap();
98 encoder.finish().unwrap();
99
100 if compressed.len() < base64.len() {
101 map.serialize_entry("zbase64", &BASE64.encode(&compressed))?;
102 } else {
103 map.serialize_entry("base64", &base64)?;
104 }
105 }
106
107 map.serialize_entry("base64", &BASE64.encode(&self.0))?;
108
109 map.end()
110 }
111}
112
113impl From<Buffer> for Vec<u8> {
114 fn from(value: Buffer) -> Self {
115 value.into_vec()
116 }
117}
118
119impl AsRef<[u8]> for Buffer {
120 fn as_ref(&self) -> &[u8] {
121 &self.0
122 }
123}
124
125impl AsMut<[u8]> for Buffer {
126 fn as_mut(&mut self) -> &mut [u8] {
127 &mut self.0
128 }
129}
130
131impl FromIterator<u8> for Buffer {
132 fn from_iter<T: IntoIterator<Item = u8>>(iter: T) -> Self {
133 Self(Vec::from_iter(iter))
134 }
135}
136
137impl Extend<u8> for Buffer {
138 fn extend<T: IntoIterator<Item = u8>>(&mut self, iter: T) {
139 self.0.extend(iter);
140 }
141}
142
143#[cfg(test)]
144mod tests {
145 use super::*;
146
147 #[test]
148 fn test_base64_de() {
149 assert_eq!(
150 serde_json::from_str::<Buffer>(
151 r#"{"m":null,"t":"buffer","base64":"aGVsbG8gd29ybGQ="}"#
152 )
153 .unwrap(),
154 Buffer::new(b"hello world".to_vec())
155 );
156 }
157
158 #[test]
159 fn test_zbase64_de() {
160 assert_eq!(
161 serde_json::from_str::<Buffer>(r#"{"m":null,"t":"buffer","zbase64":"KLUv/SBfbQAAMGhlbGxvIAEAlqkUAQ=="}"#).unwrap(),
162 Buffer::new(b"hello hello hello hello hello hello hello hello hello hello hello hello hello hello hello hello".to_vec())
163 )
164 }
165}