tachyon_web/ws/
deflate.rs1use flate2::{Compress, Compression, Decompress, FlushCompress, FlushDecompress};
12use hyper::header::{HeaderMap, HeaderValue, SEC_WEBSOCKET_EXTENSIONS};
13
14const DEFLATE_TAIL: [u8; 4] = [0x00, 0x00, 0xff, 0xff];
17
18#[derive(Debug, Clone, Copy)]
23#[non_exhaustive]
24pub struct DeflateConfig {
25 pub server_no_context_takeover: bool,
28 pub client_no_context_takeover: bool,
31 pub server_max_window_bits: u8,
34}
35
36impl Default for DeflateConfig {
37 fn default() -> Self {
38 Self {
39 server_no_context_takeover: false,
40 client_no_context_takeover: false,
41 server_max_window_bits: 15,
42 }
43 }
44}
45
46#[derive(Debug, Default, Clone, Copy)]
48pub(super) struct Offer {
49 server_no_context_takeover: bool,
50 client_no_context_takeover: bool,
51 server_max_window_bits_offered: bool,
57}
58
59#[derive(Debug, Clone, Copy)]
61pub(super) struct Agreement {
62 pub(super) server_no_context_takeover: bool,
63 pub(super) client_no_context_takeover: bool,
64 pub(super) server_max_window_bits: u8,
65 echo_server_max_window_bits: bool,
68}
69
70pub(super) fn parse_offers(headers: &HeaderMap) -> Vec<Offer> {
75 let mut offers = Vec::new();
76 for header in headers.get_all(SEC_WEBSOCKET_EXTENSIONS) {
77 let Ok(text) = header.to_str() else { continue };
78 for extension in text.split(',') {
79 let mut parts = extension.split(';').map(str::trim);
80 let Some(name) = parts.next() else { continue };
81 if !name.eq_ignore_ascii_case("permessage-deflate") {
82 continue;
83 }
84 if let Some(offer) = parse_params(parts) {
85 offers.push(offer);
86 }
87 }
88 }
89 offers
90}
91
92fn parse_window_bits(value: Option<&str>) -> Result<(), ()> {
95 match value.map(str::parse::<u8>) {
96 None | Some(Ok(9..=15)) => Ok(()),
97 Some(_) => Err(()),
98 }
99}
100
101fn parse_params<'a>(params: impl Iterator<Item = &'a str>) -> Option<Offer> {
102 let mut offer = Offer::default();
103 for param in params {
104 if param.is_empty() {
105 continue;
106 }
107 let (key, value) = match param.split_once('=') {
108 Some((k, v)) => (k.trim(), Some(v.trim().trim_matches('"'))),
109 None => (param, None),
110 };
111 match key {
112 "server_no_context_takeover" if value.is_none() => {
113 offer.server_no_context_takeover = true;
114 }
115 "client_no_context_takeover" if value.is_none() => {
116 offer.client_no_context_takeover = true;
117 }
118 "server_max_window_bits" => {
119 parse_window_bits(value).ok()?;
120 offer.server_max_window_bits_offered = true;
121 }
122 "client_max_window_bits" => {
123 parse_window_bits(value).ok()?;
124 }
125 _ => return None,
127 }
128 }
129 Some(offer)
130}
131
132pub(super) fn negotiate(offers: &[Offer], config: DeflateConfig) -> Option<Agreement> {
135 let offer = offers.first()?;
136 let server_max_window_bits = config.server_max_window_bits.clamp(9, 15);
137 Some(Agreement {
138 server_no_context_takeover: config.server_no_context_takeover
139 || offer.server_no_context_takeover,
140 client_no_context_takeover: config.client_no_context_takeover
141 || offer.client_no_context_takeover,
142 server_max_window_bits,
143 echo_server_max_window_bits: offer.server_max_window_bits_offered
144 && server_max_window_bits < 15,
145 })
146}
147
148pub(super) fn agreement_header_value(agreement: Agreement) -> HeaderValue {
150 let mut value = String::from("permessage-deflate");
151 if agreement.server_no_context_takeover {
152 value.push_str("; server_no_context_takeover");
153 }
154 if agreement.client_no_context_takeover {
155 value.push_str("; client_no_context_takeover");
156 }
157 if agreement.echo_server_max_window_bits {
158 value.push_str("; server_max_window_bits=");
159 value.push_str(&agreement.server_max_window_bits.to_string());
160 }
161 HeaderValue::from_str(&value).unwrap_or_else(|_| HeaderValue::from_static("permessage-deflate"))
162}
163
164pub(super) struct PerMessageDeflate {
166 compress: Compress,
167 decompress: Decompress,
168 server_no_context_takeover: bool,
169 client_no_context_takeover: bool,
170}
171
172impl PerMessageDeflate {
173 pub(super) fn new(agreement: Agreement) -> Self {
174 Self {
175 compress: Compress::new_with_window_bits(
176 Compression::default(),
177 false,
178 agreement.server_max_window_bits,
179 ),
180 decompress: Decompress::new_with_window_bits(false, 15),
181 server_no_context_takeover: agreement.server_no_context_takeover,
182 client_no_context_takeover: agreement.client_no_context_takeover,
183 }
184 }
185
186 pub(super) fn compress_if_smaller(
199 &mut self,
200 data: &[u8],
201 ) -> Result<Option<Vec<u8>>, crate::http::error::Error> {
202 let compressed = self.compress_raw(data)?;
203 if compressed.len() < data.len() {
204 if self.server_no_context_takeover {
205 self.compress.reset();
206 }
207 Ok(Some(compressed))
208 } else {
209 self.compress.reset();
210 Ok(None)
211 }
212 }
213
214 fn compress_raw(&mut self, data: &[u8]) -> Result<Vec<u8>, crate::http::error::Error> {
216 let total_in_before = self.compress.total_in();
217 let mut out = Vec::with_capacity(data.len() + 32);
218 loop {
219 grow(&mut out, 1024.max(data.len()));
220 self.compress
221 .compress_vec(data, &mut out, FlushCompress::Sync)
222 .map_err(|e| crate::http::error::Error::Internal(e.to_string()))?;
223 let consumed =
224 usize::try_from(self.compress.total_in() - total_in_before).unwrap_or(usize::MAX);
225 if consumed >= data.len() {
226 break;
227 }
228 }
229 out.truncate(out.len().saturating_sub(DEFLATE_TAIL.len()));
230 Ok(out)
231 }
232
233 pub(super) fn decompress(
235 &mut self,
236 data: &[u8],
237 max_size: Option<usize>,
238 ) -> Result<Vec<u8>, crate::http::error::Error> {
239 let max_size = max_size.unwrap_or(usize::MAX);
240 let mut input = Vec::with_capacity(data.len() + DEFLATE_TAIL.len());
241 input.extend_from_slice(data);
242 input.extend_from_slice(&DEFLATE_TAIL);
243
244 let total_in_before = self.decompress.total_in();
245 let mut out = Vec::with_capacity((data.len() * 3 + 32).min(max_size));
246 loop {
247 grow(&mut out, 1024);
248 let consumed_before =
249 usize::try_from(self.decompress.total_in() - total_in_before).unwrap_or(usize::MAX);
250 self.decompress
251 .decompress_vec(&input[consumed_before..], &mut out, FlushDecompress::Sync)
252 .map_err(|e| crate::http::error::Error::Internal(e.to_string()))?;
253 if out.len() > max_size {
254 return Err(crate::http::error::Error::Internal(
255 "decompressed message exceeds the configured maximum size".to_string(),
256 ));
257 }
258 let consumed =
259 usize::try_from(self.decompress.total_in() - total_in_before).unwrap_or(usize::MAX);
260 if consumed >= input.len() {
261 break;
262 }
263 }
264 if self.client_no_context_takeover {
265 self.decompress.reset(false);
266 }
267 Ok(out)
268 }
269}
270
271fn grow(out: &mut Vec<u8>, min_initial: usize) {
275 let additional = out.capacity().max(min_initial);
276 out.reserve(additional);
277}