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