Skip to main content

bynk_lsp/
transport.rs

1//! #1667: a framing pump in front of `tower-lsp`'s transport.
2//!
3//! `tower-lsp` 0.20 reads stdin through a `FramedRead`, and a framed stream
4//! ends at its first decode error. So one message whose body isn't valid JSON
5//! stopped the whole server, with exit 0 and nothing on stderr. The trigger
6//! is reachable from a well-behaved client: `JSON.stringify` writes a lone
7//! UTF-16 surrogate in a string as `\udcff`, which `serde_json` rejects.
8//!
9//! [`pump`] reads the client's frames itself and forwards to the server only
10//! bodies that parse, so a bad one can't end the stream:
11//!
12//! - a body that parses is forwarded unchanged;
13//! - a body whose only fault is a lone surrogate escape is repaired (each
14//!   lone `\uD800`–`\uDFFF` becomes `\uFFFD`) and forwarded. U+FFFD is one
15//!   UTF-16 code unit, like the surrogate it replaces, so the client's
16//!   positions into the text stay aligned with the server's;
17//! - any other body is logged and skipped. Its `id` can't be recovered from
18//!   unparseable JSON, so no error response is possible.
19//!
20//! A header block without `Content-Length`, or a stream that ends mid-frame
21//! (a truncated body), leaves nothing to resynchronise on: the pump stops,
22//! with the reason logged, and the server sees end of input.
23
24use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader};
25
26/// #1771 review: the largest body [`pump`] accepts. A `Content-Length` is
27/// otherwise trusted and allocated in full before the body arrives, so one
28/// corrupted digit could abort the process. A frame over this is unrecoverable
29/// framing, like a missing `Content-Length`. 128 MiB is far above any real
30/// message (a `didOpen` carries one file's text).
31pub const MAX_BODY: usize = 128 << 20;
32
33/// The longest header line [`pump`] reads, so a client that never sends a
34/// newline can't grow the line without limit.
35pub const MAX_HEADER_LINE: u64 = 8 << 10;
36
37/// What [`pump`] does with one frame's body.
38#[derive(Debug, PartialEq, Eq)]
39pub enum Frame {
40    /// The body parses as JSON: forward it unchanged.
41    Forward,
42    /// The body parses once its lone surrogate escapes are replaced.
43    Repaired(Vec<u8>),
44    /// The body can't be parsed. The string says why.
45    Skip(String),
46}
47
48/// Decide what to do with one frame's body.
49pub 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
63/// Replace each lone surrogate escape inside a JSON string with `\uFFFD`.
64/// `None` when there is none to replace.
65///
66/// The scan tracks string and escape state, so an escaped backslash
67/// (`\\udcff`, a literal backslash then `udcff`) is left alone, and a valid
68/// pair (`\ud83d\ude00`) is kept as written.
69fn 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                        // A high surrogate is kept only with a low one after it.
96                        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                // Any other escape: copy it and the escaped byte together, so
122                // an escaped quote or backslash can't change the string state.
123                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
141/// The four hex digits at `at`, as a UTF-16 code unit.
142fn 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
147/// Read framed messages from `input` and write the usable ones, re-framed, to
148/// `output`. Returns when `input` ends or loses its framing. Dropping `output`
149/// then ends the server's input.
150pub 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        // Headers, up to the blank line. Only `Content-Length` matters.
158        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            // The server side has gone; nothing more to deliver.
243            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        // `\\udcff` is a literal backslash then `udcff`: valid JSON as is.
278        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    /// Run the pump over `input` and return everything it forwarded.
292    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    /// #1771 review: each unrecoverable-framing path returns (the server
301    /// then sees end of input) rather than hanging, spinning or allocating.
302    #[tokio::test]
303    async fn unrecoverable_framing_ends_the_pump() {
304        // A header block with no Content-Length.
305        assert!(pumped(b"Content-Type: x\r\n\r\n{}").await.is_empty());
306        // A truncated body: 50 bytes promised, 5 sent, then end of input.
307        assert!(pumped(b"Content-Length: 50\r\n\r\nhello").await.is_empty());
308        // Only blank lines, then end of input.
309        assert!(pumped(b"\r\n\r\n\r\n").await.is_empty());
310        // A Content-Length over the limit (or past usize) allocates nothing.
311        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        // A header line with no end.
322        let endless = vec![b'x'; (MAX_HEADER_LINE as usize) * 2];
323        assert!(pumped(&endless).await.is_empty());
324    }
325
326    /// The pump forwards good frames, repairs a lone surrogate, skips an
327    /// unparseable one, and keeps going: one bad message no longer ends the
328    /// stream the server reads.
329    #[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}