1use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader};
25
26pub const MAX_BODY: usize = 128 << 20;
32
33pub const MAX_HEADER_LINE: u64 = 8 << 10;
36
37#[derive(Debug, PartialEq, Eq)]
39pub enum Frame {
40 Forward,
42 Repaired(Vec<u8>),
44 Skip(String),
46}
47
48pub fn classify(body: &[u8]) -> Frame {
50 let error = match serde_json::from_slice::<serde_json::Value>(body) {
51 Ok(_) => return Frame::Forward,
52 Err(e) => e,
53 };
54 if let Ok(text) = std::str::from_utf8(body)
55 && let Some(repaired) = repair_lone_surrogates(text)
56 && serde_json::from_str::<serde_json::Value>(&repaired).is_ok()
57 {
58 return Frame::Repaired(repaired.into_bytes());
59 }
60 Frame::Skip(error.to_string())
61}
62
63fn repair_lone_surrogates(text: &str) -> Option<String> {
70 let bytes = text.as_bytes();
71 let mut out = String::with_capacity(text.len());
72 let mut changed = false;
73 let mut in_string = false;
74 let mut i = 0;
75 while i < bytes.len() {
76 let b = bytes[i];
77 if !in_string {
78 if b == b'"' {
79 in_string = true;
80 }
81 out.push(b as char);
82 i += 1;
83 continue;
84 }
85 match b {
86 b'"' => {
87 in_string = false;
88 out.push('"');
89 i += 1;
90 }
91 b'\\' if bytes.get(i + 1) == Some(&b'u') => {
92 let unit = hex4(bytes, i + 2);
93 match unit {
94 Some(0xD800..=0xDBFF) => {
95 let low = (bytes.get(i + 6) == Some(&b'\\')
97 && bytes.get(i + 7) == Some(&b'u'))
98 .then(|| hex4(bytes, i + 8))
99 .flatten();
100 if matches!(low, Some(0xDC00..=0xDFFF)) {
101 out.push_str(&text[i..i + 12]);
102 i += 12;
103 } else {
104 out.push_str("\\uFFFD");
105 changed = true;
106 i += 6;
107 }
108 }
109 Some(0xDC00..=0xDFFF) => {
110 out.push_str("\\uFFFD");
111 changed = true;
112 i += 6;
113 }
114 _ => {
115 out.push_str("\\u");
116 i += 2;
117 }
118 }
119 }
120 b'\\' => {
121 out.push('\\');
124 if let Some(next) = text[i + 1..].chars().next() {
125 out.push(next);
126 i += 1 + next.len_utf8();
127 } else {
128 i += 1;
129 }
130 }
131 _ => {
132 let ch = text[i..].chars().next().expect("in bounds");
133 out.push(ch);
134 i += ch.len_utf8();
135 }
136 }
137 }
138 changed.then_some(out)
139}
140
141fn hex4(bytes: &[u8], at: usize) -> Option<u32> {
143 let digits = std::str::from_utf8(bytes.get(at..at + 4)?).ok()?;
144 u32::from_str_radix(digits, 16).ok()
145}
146
147pub async fn pump<R, W>(input: R, mut output: W)
151where
152 R: tokio::io::AsyncRead + Unpin,
153 W: AsyncWrite + Unpin,
154{
155 let mut input = BufReader::new(input);
156 loop {
157 let mut length: Option<usize> = None;
159 let mut saw_header = false;
160 loop {
161 let mut line = String::new();
162 let read = (&mut input)
163 .take(MAX_HEADER_LINE)
164 .read_line(&mut line)
165 .await;
166 if read.as_ref().is_ok_and(|n| *n as u64 == MAX_HEADER_LINE) && !line.ends_with('\n') {
167 tracing::error!(
168 "a client header line is over {MAX_HEADER_LINE} bytes; \
169 the input can't be resynchronised, so it is closed"
170 );
171 return;
172 }
173 match read {
174 Ok(0) => {
175 if saw_header {
176 tracing::error!("client input ended inside a message's headers");
177 }
178 return;
179 }
180 Ok(_) => {}
181 Err(e) => {
182 tracing::error!("reading client input failed: {e}");
183 return;
184 }
185 }
186 let line = line.trim_end_matches(['\r', '\n']);
187 if line.is_empty() {
188 if saw_header {
189 break;
190 }
191 continue;
192 }
193 saw_header = true;
194 if let Some((name, value)) = line.split_once(':')
195 && name.trim().eq_ignore_ascii_case("content-length")
196 {
197 length = value.trim().parse().ok();
198 }
199 }
200 let Some(length) = length else {
201 tracing::error!(
202 "a client message had no usable Content-Length header; \
203 the input can't be resynchronised, so it is closed"
204 );
205 return;
206 };
207 if length > MAX_BODY {
208 tracing::error!(
209 "a client message claims a {length}-byte body, over the {MAX_BODY}-byte limit; \
210 the input can't be resynchronised, so it is closed"
211 );
212 return;
213 }
214 let mut body = vec![0u8; length];
215 if let Err(e) = input.read_exact(&mut body).await {
216 tracing::error!(
217 "client input ended inside a message body ({length} bytes expected): {e}"
218 );
219 return;
220 }
221 let body = match classify(&body) {
222 Frame::Forward => body,
223 Frame::Repaired(repaired) => {
224 tracing::warn!(
225 "replaced lone UTF-16 surrogate escapes in a client message with U+FFFD"
226 );
227 repaired
228 }
229 Frame::Skip(reason) => {
230 let shown = String::from_utf8_lossy(&body[..body.len().min(200)]).into_owned();
231 tracing::error!(
232 "skipped a client message that isn't valid JSON ({reason}): {shown}"
233 );
234 continue;
235 }
236 };
237 let header = format!("Content-Length: {}\r\n\r\n", body.len());
238 if output.write_all(header.as_bytes()).await.is_err()
239 || output.write_all(&body).await.is_err()
240 || output.flush().await.is_err()
241 {
242 return;
244 }
245 }
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251
252 #[test]
253 fn valid_json_is_forwarded() {
254 assert_eq!(classify(br#"{"a":"\ud83d\ude00"}"#), Frame::Forward);
255 }
256
257 #[test]
258 fn a_lone_surrogate_is_repaired_to_the_replacement_character() {
259 let Frame::Repaired(body) = classify(br#"{"text":"a\udcffb"}"#) else {
260 panic!("repaired");
261 };
262 let value: serde_json::Value = serde_json::from_slice(&body).unwrap();
263 assert_eq!(value["text"], "a\u{FFFD}b");
264 }
265
266 #[test]
267 fn a_lone_high_surrogate_is_repaired_and_a_pair_is_kept() {
268 let Frame::Repaired(body) = classify(br#"{"t":"\ud83d x \ud83d\ude00"}"#) else {
269 panic!("repaired");
270 };
271 let value: serde_json::Value = serde_json::from_slice(&body).unwrap();
272 assert_eq!(value["t"], "\u{FFFD} x \u{1F600}");
273 }
274
275 #[test]
276 fn an_escaped_backslash_before_u_is_not_a_surrogate_escape() {
277 assert_eq!(classify(br#"{"t":"\\udcff"}"#), Frame::Forward);
279 }
280
281 #[test]
282 fn other_invalid_json_is_skipped() {
283 assert!(matches!(classify(b"{not json"), Frame::Skip(_)));
284 }
285
286 #[test]
287 fn a_surrogate_escape_cut_off_at_the_end_is_skipped() {
288 assert!(matches!(classify(br#"{"t":"\ud8"#), Frame::Skip(_)));
289 }
290
291 async fn pumped(input: &[u8]) -> Vec<u8> {
293 let (mut server_side, client_side) = tokio::io::duplex(64 * 1024);
294 pump(input, client_side).await;
295 let mut out = Vec::new();
296 server_side.read_to_end(&mut out).await.unwrap();
297 out
298 }
299
300 #[tokio::test]
303 async fn unrecoverable_framing_ends_the_pump() {
304 assert!(pumped(b"Content-Type: x\r\n\r\n{}").await.is_empty());
306 assert!(pumped(b"Content-Length: 50\r\n\r\nhello").await.is_empty());
308 assert!(pumped(b"\r\n\r\n\r\n").await.is_empty());
310 assert!(
312 pumped(b"Content-Length: 99999999999\r\n\r\n{}")
313 .await
314 .is_empty()
315 );
316 assert!(
317 pumped(b"Content-Length: 18446744073709551615\r\n\r\n{}")
318 .await
319 .is_empty()
320 );
321 let endless = vec![b'x'; (MAX_HEADER_LINE as usize) * 2];
323 assert!(pumped(&endless).await.is_empty());
324 }
325
326 #[tokio::test]
330 async fn the_pump_skips_a_bad_frame_and_keeps_going() {
331 let frame = |body: &str| format!("Content-Length: {}\r\n\r\n{body}", body.len());
332 let input = [
333 frame(r#"{"n":1}"#),
334 frame("{not json"),
335 frame(r#"{"t":"\udcff"}"#),
336 frame(r#"{"n":2}"#),
337 ]
338 .concat();
339 let (server_side, client_side) = tokio::io::duplex(64 * 1024);
340 pump(input.as_bytes(), client_side).await;
341 let mut out = String::new();
342 let mut server_side = server_side;
343 server_side.read_to_string(&mut out).await.unwrap();
344 assert_eq!(out.matches("Content-Length").count(), 3, "{out}");
345 assert!(
346 out.contains(r#"{"n":1}"#) && out.contains(r#"{"n":2}"#),
347 "{out}"
348 );
349 assert!(out.contains(r#"\uFFFD"#), "{out}");
350 assert!(!out.contains("not json"), "{out}");
351 }
352}