strop_ui_protocol/
frame.rs1use std::io::{self, Read, Write};
9
10pub const MAX_HEADER_BYTES: usize = 8192;
13
14pub const MAX_BODY_BYTES: usize = 32 * 1024 * 1024;
17
18#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
22pub enum FrameError {
23 #[error("frame header exceeded the {limit}-byte bound without a terminator")]
24 HeaderTooLarge { limit: usize },
25 #[error("frame body length {actual} exceeds the {limit}-byte bound")]
26 BodyTooLarge { limit: usize, actual: usize },
27 #[error("frame header carries no valid Content-Length")]
28 MissingLength,
29 #[error("frame header line is not `Name: value`")]
30 MalformedHeader,
31 #[error("stream ended mid-frame")]
33 Truncated,
34}
35
36pub fn write_frame(mut writer: impl Write, body: &[u8]) -> io::Result<()> {
38 write!(writer, "Content-Length: {}\r\n\r\n", body.len())?;
39 writer.write_all(body)?;
40 writer.flush()
41}
42
43pub fn write_message(writer: impl Write, message: &impl serde::Serialize) -> io::Result<()> {
45 write_frame(writer, &serde_json::to_vec(message)?)
46}
47
48#[derive(Debug, Default)]
52pub struct FrameDecoder {
53 buffer: Vec<u8>,
54}
55
56impl FrameDecoder {
57 pub fn new() -> Self {
58 Self::default()
59 }
60
61 pub fn is_empty(&self) -> bool {
63 self.buffer.is_empty()
64 }
65
66 pub fn accept(&mut self, bytes: &[u8]) -> Result<(), FrameError> {
69 self.buffer.extend_from_slice(bytes);
70 if !self.header_complete() && self.buffer.len() > MAX_HEADER_BYTES {
71 return Err(FrameError::HeaderTooLarge {
72 limit: MAX_HEADER_BYTES,
73 });
74 }
75 Ok(())
76 }
77
78 pub fn next_frame(&mut self) -> Result<Option<Vec<u8>>, FrameError> {
80 let Some(header_end) = find(&self.buffer, b"\r\n\r\n") else {
81 return Ok(None);
82 };
83 if header_end + 4 > MAX_HEADER_BYTES {
84 return Err(FrameError::HeaderTooLarge {
85 limit: MAX_HEADER_BYTES,
86 });
87 }
88 let length = content_length(&self.buffer[..header_end])?;
89 if length > MAX_BODY_BYTES {
90 return Err(FrameError::BodyTooLarge {
91 limit: MAX_BODY_BYTES,
92 actual: length,
93 });
94 }
95 let start = header_end + 4;
96 if self.buffer.len() < start + length {
97 return Ok(None);
98 }
99 let body = self.buffer[start..start + length].to_vec();
100 self.buffer.drain(..start + length);
101 Ok(Some(body))
102 }
103
104 fn header_complete(&self) -> bool {
105 find(&self.buffer, b"\r\n\r\n").is_some()
106 }
107}
108
109pub fn read_frame(
114 reader: &mut impl Read,
115 decoder: &mut FrameDecoder,
116) -> io::Result<Option<Vec<u8>>> {
117 let mut chunk = [0u8; 8192];
118 loop {
119 match decoder.next_frame() {
120 Ok(Some(body)) => return Ok(Some(body)),
121 Ok(None) => {}
122 Err(error) => return Err(io::Error::new(io::ErrorKind::InvalidData, error)),
123 }
124 let read = reader.read(&mut chunk)?;
125 if read == 0 {
126 return if decoder.is_empty() {
127 Ok(None)
128 } else {
129 Err(io::Error::new(
130 io::ErrorKind::UnexpectedEof,
131 "stream ended mid-frame",
132 ))
133 };
134 }
135 decoder
136 .accept(&chunk[..read])
137 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
138 }
139}
140
141fn find(haystack: &[u8], needle: &[u8]) -> Option<usize> {
142 haystack
143 .windows(needle.len())
144 .position(|window| window == needle)
145}
146
147fn content_length(header: &[u8]) -> Result<usize, FrameError> {
148 let header = std::str::from_utf8(header).map_err(|_| FrameError::MalformedHeader)?;
149 let mut length = None;
150 for line in header.split("\r\n") {
151 let (name, value) = line.split_once(':').ok_or(FrameError::MalformedHeader)?;
152 if name.trim().is_empty() {
153 return Err(FrameError::MalformedHeader);
154 }
155 if name.trim() == "Content-Length" {
156 length = Some(
157 value
158 .trim()
159 .parse::<usize>()
160 .map_err(|_| FrameError::MissingLength)?,
161 );
162 }
163 }
164 length.ok_or(FrameError::MissingLength)
165}
166
167#[cfg(test)]
168mod tests {
169 use super::*;
170
171 fn frame(body: &[u8]) -> Vec<u8> {
172 let mut bytes = Vec::new();
173 write_frame(&mut bytes, body).unwrap();
174 bytes
175 }
176
177 #[test]
178 fn roundtrip_whole_and_concatenated() {
179 let mut bytes = frame(b"first");
180 bytes.extend(frame(b"second"));
181 let mut decoder = FrameDecoder::new();
182 decoder.accept(&bytes).unwrap();
183 assert_eq!(
184 decoder.next_frame().unwrap().as_deref(),
185 Some(&b"first"[..])
186 );
187 assert_eq!(
188 decoder.next_frame().unwrap().as_deref(),
189 Some(&b"second"[..])
190 );
191 assert!(decoder.next_frame().unwrap().is_none());
192 assert!(decoder.is_empty());
193 }
194
195 #[test]
196 fn partial_io_reassembles_byte_by_byte() {
197 let bytes = frame(b"{\"hello\":\"world\"}");
198 let mut decoder = FrameDecoder::new();
199 let mut frames = Vec::new();
200 for byte in bytes {
201 decoder.accept(&[byte]).unwrap();
202 while let Some(body) = decoder.next_frame().unwrap() {
203 frames.push(body);
204 }
205 }
206 assert_eq!(frames, vec![b"{\"hello\":\"world\"}".to_vec()]);
207 }
208
209 #[test]
210 fn extra_headers_are_ignored_like_lsp() {
211 let mut bytes = b"Content-Type: application/vscode-jsonrpc; charset=utf-8\r\n".to_vec();
212 bytes.extend(frame(b"body"));
213 let mut decoder = FrameDecoder::new();
215 decoder.accept(&bytes).unwrap();
216 assert_eq!(decoder.next_frame().unwrap().as_deref(), Some(&b"body"[..]));
217 }
218
219 #[test]
220 fn oversized_header_is_typed_before_unbounded_buffering() {
221 let mut decoder = FrameDecoder::new();
222 let error = decoder
223 .accept(&vec![b'x'; MAX_HEADER_BYTES + 1])
224 .unwrap_err();
225 assert_eq!(
226 error,
227 FrameError::HeaderTooLarge {
228 limit: MAX_HEADER_BYTES
229 }
230 );
231 }
232
233 #[test]
234 fn oversized_body_is_typed_at_the_header() {
235 let declared = MAX_BODY_BYTES + 1;
236 let mut decoder = FrameDecoder::new();
237 decoder
238 .accept(format!("Content-Length: {declared}\r\n\r\n").as_bytes())
239 .unwrap();
240 assert_eq!(
241 decoder.next_frame().unwrap_err(),
242 FrameError::BodyTooLarge {
243 limit: MAX_BODY_BYTES,
244 actual: declared
245 }
246 );
247 }
248
249 #[test]
250 fn missing_and_malformed_headers_are_typed() {
251 let mut decoder = FrameDecoder::new();
252 decoder.accept(b"Content-Type: text\r\n\r\n").unwrap();
253 assert_eq!(decoder.next_frame().unwrap_err(), FrameError::MissingLength);
254
255 let mut decoder = FrameDecoder::new();
256 decoder.accept(b"no-colon-here\r\n\r\n").unwrap();
257 assert_eq!(
258 decoder.next_frame().unwrap_err(),
259 FrameError::MalformedHeader
260 );
261
262 let mut decoder = FrameDecoder::new();
263 decoder.accept(b"Content-Length: many\r\n\r\n").unwrap();
264 assert_eq!(decoder.next_frame().unwrap_err(), FrameError::MissingLength);
265 }
266
267 #[test]
268 fn blocking_read_distinguishes_clean_eof_from_truncation() {
269 let mut bytes = frame(b"done");
270 let mut decoder = FrameDecoder::new();
271 let mut cursor = io::Cursor::new(bytes.clone());
272 assert_eq!(
273 read_frame(&mut cursor, &mut decoder).unwrap().as_deref(),
274 Some(&b"done"[..])
275 );
276 assert!(read_frame(&mut cursor, &mut decoder).unwrap().is_none());
277
278 bytes.truncate(bytes.len() - 2); let mut decoder = FrameDecoder::new();
280 let mut cursor = io::Cursor::new(bytes);
281 let error = read_frame(&mut cursor, &mut decoder).unwrap_err();
282 assert_eq!(error.kind(), io::ErrorKind::UnexpectedEof);
283 }
284}