1use thiserror::Error;
31use tokio::io::{AsyncBufRead, AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
32
33use crate::framing::{self, FrameError, FrameHeader};
34use crate::jsonrpc::{Message, MessageDecodeError};
35
36pub struct LspReader<R> {
43 inner: R,
44 line_buf: Vec<u8>,
48 header_buf: Vec<u8>,
50 body_buf: Vec<u8>,
53}
54
55impl<R: AsyncRead + Unpin> LspReader<BufReader<R>> {
56 pub fn new(inner: R) -> Self {
60 Self::from_buf_reader(BufReader::new(inner))
61 }
62}
63
64impl<R: AsyncBufRead + Unpin> LspReader<R> {
65 pub fn from_buf_reader(inner: R) -> Self {
69 Self {
70 inner,
71 line_buf: Vec::with_capacity(64),
72 header_buf: Vec::with_capacity(128),
73 body_buf: Vec::with_capacity(1024),
74 }
75 }
76
77 pub async fn read_message(&mut self) -> Result<Option<Message>, CodecError> {
84 self.header_buf.clear();
87 let mut got_any = false;
88 loop {
89 self.line_buf.clear();
90 let n = self
91 .inner
92 .read_until(b'\n', &mut self.line_buf)
93 .await
94 .map_err(CodecError::Io)?;
95 if n == 0 {
96 if !got_any {
98 return Ok(None);
99 }
100 return Err(CodecError::UnexpectedEof);
101 }
102 got_any = true;
103
104 let line = strip_crlf(&self.line_buf);
106
107 if line.is_empty() {
108 break;
110 }
111 if !self.header_buf.is_empty() {
112 self.header_buf.extend_from_slice(b"\r\n");
113 }
114 self.header_buf.extend_from_slice(line);
115 }
116
117 let header: FrameHeader =
118 framing::parse_header_block(&self.header_buf).map_err(CodecError::Frame)?;
119
120 let len = header.content_length as usize;
121 self.body_buf.clear();
122 self.body_buf.resize(len, 0);
123 tokio::io::AsyncReadExt::read_exact(&mut self.inner, &mut self.body_buf)
124 .await
125 .map_err(CodecError::Io)?;
126
127 let msg = Message::from_json(&self.body_buf).map_err(CodecError::Decode)?;
128 Ok(Some(msg))
129 }
130}
131
132pub struct LspWriter<W> {
137 inner: W,
138 encode_buf: Vec<u8>,
140}
141
142impl<W: AsyncWrite + Unpin> LspWriter<W> {
143 pub fn new(inner: W) -> Self {
144 Self {
145 inner,
146 encode_buf: Vec::with_capacity(1024),
147 }
148 }
149
150 pub async fn write_message(&mut self, msg: &Message) -> Result<(), CodecError> {
155 self.encode_buf.clear();
156 let body = msg
157 .to_json()
158 .map_err(|e| CodecError::Decode(MessageDecodeError::Json(e)))?;
159 let header = framing::encode_header(body.len());
160 self.inner
166 .write_all(header.as_bytes())
167 .await
168 .map_err(CodecError::Io)?;
169 self.inner.write_all(&body).await.map_err(CodecError::Io)?;
170 self.inner.flush().await.map_err(CodecError::Io)?;
171 Ok(())
172 }
173
174 pub fn inner_mut(&mut self) -> &mut W {
177 &mut self.inner
178 }
179}
180
181#[derive(Debug, Error)]
183pub enum CodecError {
184 #[error("io error: {0}")]
186 Io(#[source] std::io::Error),
187 #[error("unexpected EOF mid-message")]
190 UnexpectedEof,
191 #[error("framing: {0}")]
193 Frame(#[source] FrameError),
194 #[error("decode: {0}")]
197 Decode(#[source] MessageDecodeError),
198}
199
200fn strip_crlf(line: &[u8]) -> &[u8] {
201 let mut end = line.len();
202 if end > 0 && line[end - 1] == b'\n' {
203 end -= 1;
204 }
205 if end > 0 && line[end - 1] == b'\r' {
206 end -= 1;
207 }
208 &line[..end]
209}
210
211#[cfg(test)]
212mod tests {
213 use super::*;
214 use crate::jsonrpc::{Notification, Request, RequestId, Response};
215 use serde_json::json;
216 use tokio::io::duplex;
217
218 fn frame_one(body: &[u8]) -> Vec<u8> {
220 let mut out = format!("Content-Length: {}\r\n\r\n", body.len()).into_bytes();
221 out.extend_from_slice(body);
222 out
223 }
224
225 #[tokio::test]
226 async fn reads_one_request() {
227 let body = br#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}"#;
228 let stream = frame_one(body);
229 let mut r = LspReader::new(&stream[..]);
230 let msg = r.read_message().await.unwrap().unwrap();
231 match msg {
232 Message::Request(req) => {
233 assert_eq!(req.method, "initialize");
234 assert_eq!(req.id, RequestId::Number(1));
235 }
236 _ => panic!("expected request"),
237 }
238 }
239
240 #[tokio::test]
241 async fn reads_back_to_back_messages() {
242 let mut stream = Vec::new();
243 stream.extend_from_slice(&frame_one(
244 br#"{"jsonrpc":"2.0","id":1,"method":"initialize"}"#,
245 ));
246 stream.extend_from_slice(&frame_one(
247 br#"{"jsonrpc":"2.0","method":"initialized","params":{}}"#,
248 ));
249 stream.extend_from_slice(&frame_one(
250 br#"{"jsonrpc":"2.0","id":2,"method":"shutdown"}"#,
251 ));
252 let mut r = LspReader::new(&stream[..]);
253 assert!(matches!(
255 r.read_message().await.unwrap().unwrap(),
256 Message::Request(_)
257 ));
258 assert!(matches!(
259 r.read_message().await.unwrap().unwrap(),
260 Message::Notification(_)
261 ));
262 assert!(matches!(
263 r.read_message().await.unwrap().unwrap(),
264 Message::Request(_)
265 ));
266 assert!(r.read_message().await.unwrap().is_none());
267 }
268
269 #[tokio::test]
270 async fn clean_eof_before_first_message_returns_none() {
271 let mut r = LspReader::new(&[][..]);
272 assert!(r.read_message().await.unwrap().is_none());
273 }
274
275 #[tokio::test]
276 async fn eof_mid_header_is_error() {
277 let stream = b"Content-Length: 5\r\n";
278 let mut r = LspReader::new(&stream[..]);
279 let err = r.read_message().await.unwrap_err();
280 assert!(matches!(err, CodecError::UnexpectedEof));
281 }
282
283 #[tokio::test]
284 async fn eof_mid_body_is_error() {
285 let mut stream = b"Content-Length: 100\r\n\r\n".to_vec();
287 stream.extend_from_slice(b"hello");
288 let mut r = LspReader::new(&stream[..]);
289 let err = r.read_message().await.unwrap_err();
290 assert!(matches!(err, CodecError::Io(_)));
291 }
292
293 #[tokio::test]
294 async fn malformed_header_is_error() {
295 let stream = b"Content-Length: not-a-number\r\n\r\n";
296 let mut r = LspReader::new(&stream[..]);
297 let err = r.read_message().await.unwrap_err();
298 assert!(matches!(err, CodecError::Frame(_)));
299 }
300
301 #[tokio::test]
302 async fn invalid_json_body_is_error() {
303 let body = b"{ not json";
304 let stream = frame_one(body);
305 let mut r = LspReader::new(&stream[..]);
306 let err = r.read_message().await.unwrap_err();
307 assert!(matches!(err, CodecError::Decode(_)));
308 }
309
310 #[tokio::test]
311 async fn write_message_produces_valid_frame() {
312 let mut buf = Vec::new();
313 {
314 let mut w = LspWriter::new(&mut buf);
315 let req = Message::Request(Request::new(
316 RequestId::from_u64(5),
317 "textDocument/hover",
318 Some(json!({"textDocument": {"uri": "file:///a.rs"}})),
319 ));
320 w.write_message(&req).await.unwrap();
321 }
322 let mut r = LspReader::new(&buf[..]);
324 let msg = r.read_message().await.unwrap().unwrap();
325 match msg {
326 Message::Request(req) => {
327 assert_eq!(req.method, "textDocument/hover");
328 assert_eq!(req.id, RequestId::Number(5));
329 }
330 _ => panic!("expected request"),
331 }
332 }
333
334 #[tokio::test]
335 async fn write_then_read_via_duplex_pipe() {
336 let (a, b) = duplex(64 * 1024);
340 let (a_read, a_write) = tokio::io::split(a);
341 let (b_read, b_write) = tokio::io::split(b);
342
343 let mut writer = LspWriter::new(a_write);
344 let mut reader = LspReader::new(b_read);
345 let _bridge_read = a_read;
346 let _bridge_write = b_write;
347
348 let n = Message::Notification(Notification::new(
349 "$/progress",
350 Some(json!({"token": "t1"})),
351 ));
352 writer.write_message(&n).await.unwrap();
353 let got = reader.read_message().await.unwrap().unwrap();
354 assert!(matches!(got, Message::Notification(_)));
355 }
356
357 #[tokio::test]
358 async fn handles_lf_only_line_endings() {
359 let body = br#"{"jsonrpc":"2.0","id":1,"method":"x"}"#;
363 let mut stream = format!("Content-Length: {}\n\n", body.len()).into_bytes();
364 stream.extend_from_slice(body);
365 let mut r = LspReader::new(&stream[..]);
366 let msg = r.read_message().await.unwrap().unwrap();
367 assert!(matches!(msg, Message::Request(_)));
368 }
369
370 #[tokio::test]
371 async fn round_trip_response_through_pipe() {
372 let mut buf = Vec::new();
373 {
374 let mut w = LspWriter::new(&mut buf);
375 let resp = Message::Response(Response::ok(
376 RequestId::from_u64(1),
377 json!({"capabilities": {"hoverProvider": true}}),
378 ));
379 w.write_message(&resp).await.unwrap();
380 }
381 let mut r = LspReader::new(&buf[..]);
382 let msg = r.read_message().await.unwrap().unwrap();
383 match msg {
384 Message::Response(resp) => {
385 assert!(resp.error.is_none());
386 assert_eq!(resp.id, RequestId::Number(1));
387 }
388 _ => panic!("expected response"),
389 }
390 }
391}