1use std::io::{self, BufRead, Write};
4
5pub const MAX_MESSAGE_BYTES: usize = 8 * 1024 * 1024;
7pub const MAX_HEADER_LINES: usize = 32;
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq)]
12pub enum Framing {
13 Jsonl,
14 Headers,
15}
16
17#[derive(Debug)]
19pub enum FrameError {
20 Io(io::Error),
21 TooManyHeaders,
22 MissingContentLength,
23 InvalidContentLength,
24 MessageTooLarge,
25 Incomplete,
26}
27
28impl std::fmt::Display for FrameError {
29 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
30 match self {
31 FrameError::Io(err) => write!(f, "{err}"),
32 FrameError::TooManyHeaders => write!(f, "too many framing headers"),
33 FrameError::MissingContentLength => write!(f, "missing Content-Length"),
34 FrameError::InvalidContentLength => write!(f, "invalid Content-Length"),
35 FrameError::MessageTooLarge => write!(f, "message exceeds size limit"),
36 FrameError::Incomplete => write!(f, "incomplete frame"),
37 }
38 }
39}
40
41impl std::error::Error for FrameError {
42 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
43 match self {
44 FrameError::Io(err) => Some(err),
45 _ => None,
46 }
47 }
48}
49
50impl From<io::Error> for FrameError {
51 fn from(err: io::Error) -> Self {
52 if err.kind() == io::ErrorKind::UnexpectedEof {
53 FrameError::Incomplete
54 } else {
55 FrameError::Io(err)
56 }
57 }
58}
59
60pub fn is_header_line(line: &[u8]) -> bool {
63 let line = trim_crlf(line);
64 let mut bytes = line.iter();
65 match bytes.next() {
66 Some(b) if b.is_ascii_alphabetic() => {}
67 _ => return false,
68 }
69 for b in bytes {
70 if *b == b':' {
71 return true;
72 }
73 if !b.is_ascii_alphanumeric() && *b != b'-' {
74 return false;
75 }
76 }
77 false
78}
79
80pub fn classify_first_line(line: &[u8]) -> Framing {
82 if is_header_line(line) {
83 Framing::Headers
84 } else {
85 Framing::Jsonl
86 }
87}
88
89pub fn write_message<W: Write>(writer: &mut W, payload: &[u8], framing: Framing) -> io::Result<()> {
91 if payload.len() > MAX_MESSAGE_BYTES {
92 return Err(io::Error::new(
93 io::ErrorKind::InvalidInput,
94 "message exceeds size limit",
95 ));
96 }
97 match framing {
98 Framing::Jsonl => {
99 writer.write_all(payload)?;
100 writer.write_all(b"\n")?;
101 }
102 Framing::Headers => {
103 let header = format!("Content-Length: {}\r\n\r\n", payload.len());
104 writer.write_all(header.as_bytes())?;
105 writer.write_all(payload)?;
106 }
107 }
108 Ok(())
109}
110
111pub fn read_message<R: BufRead>(reader: &mut R) -> Result<(Vec<u8>, Framing), FrameError> {
116 let first = read_line_limited(reader, MAX_MESSAGE_BYTES + 2)?;
117 if is_header_line(&first) {
118 let length = read_content_length(reader, first)?;
119 let mut body = vec![0u8; length];
120 reader.read_exact(&mut body)?;
121 Ok((body, Framing::Headers))
122 } else {
123 let body = strip_crlf(first);
124 if body.len() > MAX_MESSAGE_BYTES {
125 return Err(FrameError::MessageTooLarge);
126 }
127 Ok((body, Framing::Jsonl))
128 }
129}
130
131fn read_content_length<R: BufRead>(reader: &mut R, first: Vec<u8>) -> Result<usize, FrameError> {
132 let mut header = first;
133 let mut count = 0;
134 let mut length = None;
135 loop {
136 if is_blank_line(&header) {
137 break;
138 }
139 count += 1;
140 if count > MAX_HEADER_LINES {
141 return Err(FrameError::TooManyHeaders);
142 }
143 if let Some(parsed) = parse_content_length_line(&header) {
144 length = Some(parsed?);
145 }
146 header = read_line_limited(reader, MAX_MESSAGE_BYTES + 1)?;
147 }
148 match length {
149 None => Err(FrameError::MissingContentLength),
150 Some(n) if n > MAX_MESSAGE_BYTES => Err(FrameError::MessageTooLarge),
151 Some(n) => Ok(n),
152 }
153}
154
155fn parse_content_length_line(line: &[u8]) -> Option<Result<usize, FrameError>> {
156 let line = trim_crlf(line);
157 let colon = line.iter().position(|&b| b == b':')?;
158 let name = trim_ascii(&line[..colon]);
159 if !name.eq_ignore_ascii_case(b"content-length") {
160 return None;
161 }
162 let value = trim_ascii(&line[colon + 1..]);
163 let text = match std::str::from_utf8(value) {
164 Ok(s) => s,
165 Err(_) => return Some(Err(FrameError::InvalidContentLength)),
166 };
167 match text.parse::<usize>() {
168 Ok(n) => Some(Ok(n)),
169 Err(_) => Some(Err(FrameError::InvalidContentLength)),
170 }
171}
172
173fn read_line_limited<R: BufRead>(reader: &mut R, max: usize) -> Result<Vec<u8>, FrameError> {
174 let mut buf = Vec::new();
175 loop {
176 let avail = reader.fill_buf()?;
177 if avail.is_empty() {
178 if buf.is_empty() {
179 return Err(FrameError::Incomplete);
180 }
181 return Ok(buf);
182 }
183 if let Some(pos) = avail.iter().position(|&b| b == b'\n') {
184 let take = pos + 1;
185 if buf.len() + take > max {
186 return Err(FrameError::MessageTooLarge);
187 }
188 buf.extend_from_slice(&avail[..take]);
189 reader.consume(take);
190 return Ok(buf);
191 }
192 if buf.len() + avail.len() > max {
193 return Err(FrameError::MessageTooLarge);
194 }
195 buf.extend_from_slice(avail);
196 let n = avail.len();
197 reader.consume(n);
198 }
199}
200
201fn is_blank_line(line: &[u8]) -> bool {
202 matches!(trim_crlf(line), b"")
203}
204
205fn strip_crlf(mut line: Vec<u8>) -> Vec<u8> {
206 if line.last() == Some(&b'\n') {
207 line.pop();
208 if line.last() == Some(&b'\r') {
209 line.pop();
210 }
211 }
212 line
213}
214
215fn trim_crlf(line: &[u8]) -> &[u8] {
216 let mut end = line.len();
217 if end > 0 && line[end - 1] == b'\n' {
218 end -= 1;
219 if end > 0 && line[end - 1] == b'\r' {
220 end -= 1;
221 }
222 }
223 &line[..end]
224}
225
226fn trim_ascii(bytes: &[u8]) -> &[u8] {
227 let start = bytes
228 .iter()
229 .position(|b| !b.is_ascii_whitespace())
230 .unwrap_or(bytes.len());
231 let end = bytes
232 .iter()
233 .rposition(|b| !b.is_ascii_whitespace())
234 .map(|i| i + 1)
235 .unwrap_or(0);
236 if start >= end {
237 &[]
238 } else {
239 &bytes[start..end]
240 }
241}
242
243#[cfg(test)]
244mod tests {
245 use super::*;
246 use std::io::Cursor;
247
248 fn encode_decode(payload: &[u8], framing: Framing) -> (Vec<u8>, Framing) {
249 let mut buf = Vec::new();
250 write_message(&mut buf, payload, framing).unwrap();
251 let mut cur = Cursor::new(buf);
252 read_message(&mut cur).unwrap()
253 }
254
255 #[test]
256 fn jsonl_request_roundtrips() {
257 let payload = br#"{"jsonrpc":"2.0","id":1,"method":"initialize"}"#;
258 let (got, framing) = encode_decode(payload, Framing::Jsonl);
259 assert_eq!(framing, Framing::Jsonl);
260 assert_eq!(got, payload);
261 }
262
263 #[test]
264 fn content_length_request_roundtrips() {
265 let payload = br#"{"jsonrpc":"2.0","id":1,"method":"identity/get"}"#;
266 let (got, framing) = encode_decode(payload, Framing::Headers);
267 assert_eq!(framing, Framing::Headers);
268 assert_eq!(got, payload);
269 }
270
271 #[test]
272 fn first_line_that_is_not_brace_is_still_jsonl() {
273 let payload = br#""not-an-object""#;
274 assert_eq!(classify_first_line(payload), Framing::Jsonl);
275 assert_eq!(classify_first_line(b"[1,2,3]\n"), Framing::Jsonl);
276 assert_eq!(classify_first_line(b"\n"), Framing::Jsonl);
277 assert_eq!(
278 classify_first_line(b"\xef\xbb\xbf{\"a\":1}"),
279 Framing::Jsonl
280 );
281 let (got, framing) = encode_decode(payload, Framing::Jsonl);
282 assert_eq!(framing, Framing::Jsonl);
283 assert_eq!(got, payload);
284 }
285
286 #[test]
287 fn header_line_is_classified_as_headers() {
288 assert_eq!(
289 classify_first_line(b"Content-Length: 12\r\n"),
290 Framing::Headers
291 );
292 assert_eq!(
293 classify_first_line(b"Content-Type: application/vscode-jsonrpc; charset=utf-8\r\n"),
294 Framing::Headers
295 );
296 assert!(is_header_line(b"Content-Length: 1"));
297 assert!(!is_header_line(b"{"));
298 assert!(!is_header_line(b"1: not a header name start"));
299 }
300
301 #[test]
302 fn headers_accept_content_type_before_length() {
303 let payload = br#"{"ok":true}"#;
304 let mut buf = Vec::new();
305 buf.extend_from_slice(b"Content-Type: application/vscode-jsonrpc\r\n");
306 buf.extend_from_slice(format!("Content-Length: {}\r\n\r\n", payload.len()).as_bytes());
307 buf.extend_from_slice(payload);
308 let (got, framing) = read_message(&mut Cursor::new(buf)).unwrap();
309 assert_eq!(framing, Framing::Headers);
310 assert_eq!(got, payload);
311 }
312
313 #[test]
314 fn missing_content_length_is_an_error() {
315 let buf = b"Content-Type: application/json\r\n\r\n";
316 let err = read_message(&mut Cursor::new(&buf[..])).unwrap_err();
317 assert!(matches!(err, FrameError::MissingContentLength));
318 }
319
320 #[test]
321 fn invalid_content_length_is_an_error() {
322 let buf = b"Content-Length: nope\r\n\r\n";
323 let err = read_message(&mut Cursor::new(&buf[..])).unwrap_err();
324 assert!(matches!(err, FrameError::InvalidContentLength));
325 }
326
327 #[test]
328 fn too_many_headers_is_an_error() {
329 let mut buf = Vec::new();
330 for _ in 0..(MAX_HEADER_LINES + 1) {
331 buf.extend_from_slice(b"X-Extra: 1\r\n");
332 }
333 buf.extend_from_slice(b"\r\n");
334 let err = read_message(&mut Cursor::new(buf)).unwrap_err();
335 assert!(matches!(err, FrameError::TooManyHeaders));
336 }
337
338 #[test]
339 fn oversized_jsonl_is_rejected() {
340 let mut line = vec![b'x'; MAX_MESSAGE_BYTES + 2];
341 line.push(b'\n');
342 let err = read_message(&mut Cursor::new(line)).unwrap_err();
343 assert!(matches!(err, FrameError::MessageTooLarge));
344 }
345
346 #[test]
347 fn oversized_content_length_is_rejected() {
348 let buf = format!("Content-Length: {}\r\n\r\n", MAX_MESSAGE_BYTES + 1);
349 let err = read_message(&mut Cursor::new(buf.into_bytes())).unwrap_err();
350 assert!(matches!(err, FrameError::MessageTooLarge));
351 }
352
353 #[test]
354 fn empty_reader_is_incomplete() {
355 let err = read_message(&mut Cursor::new(&b""[..])).unwrap_err();
356 assert!(matches!(err, FrameError::Incomplete));
357 }
358
359 #[test]
360 fn exact_max_payload_roundtrips_jsonl_and_headers() {
361 let payload = vec![b'a'; MAX_MESSAGE_BYTES];
362 let (got, framing) = encode_decode(&payload, Framing::Jsonl);
363 assert_eq!(framing, Framing::Jsonl);
364 assert_eq!(got, payload);
365 let (got, framing) = encode_decode(&payload, Framing::Headers);
366 assert_eq!(framing, Framing::Headers);
367 assert_eq!(got, payload);
368 }
369
370 #[test]
371 fn write_rejects_oversized_payload() {
372 let payload = vec![b'a'; MAX_MESSAGE_BYTES + 1];
373 let err = write_message(&mut Vec::new(), &payload, Framing::Jsonl).unwrap_err();
374 assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
375 }
376
377 #[test]
378 fn frame_error_display_names_the_limit() {
379 assert_eq!(
380 FrameError::MessageTooLarge.to_string(),
381 "message exceeds size limit"
382 );
383 assert_eq!(
384 FrameError::TooManyHeaders.to_string(),
385 "too many framing headers"
386 );
387 }
388}