1use crate::language_catalog::LanguageId;
2use lsp_types::Uri;
3use serde::de::DeserializeOwned;
4use serde::{Deserialize, Serialize};
5use serde_json::Value;
6use std::io;
7use std::marker::PhantomData;
8use std::path::PathBuf;
9use tokio_util::bytes::{Bytes, BytesMut};
10use tokio_util::codec::{Decoder, Encoder, FramedRead, FramedWrite, LengthDelimitedCodec};
11
12#[doc = include_str!("docs/protocol.md")]
13#[derive(Debug, Clone, Serialize, Deserialize)]
14pub enum DaemonRequest {
15 Initialize(InitializeRequest),
16 LspCall {
17 client_id: i64,
18 method: String,
19 params: Value,
20 },
21 GetDiagnostics {
22 client_id: i64,
23 uri: Option<Uri>,
25 },
26 QueueDiagnosticRefresh {
27 client_id: i64,
28 uri: Uri,
29 },
30 Disconnect,
31 Ping,
32}
33
34#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct InitializeRequest {
37 pub workspace_root: PathBuf,
38 pub language: LanguageId,
39}
40
41#[derive(Debug, Clone, Serialize, Deserialize)]
43pub struct LspNotification {
44 pub method: String,
45 pub params: Value,
46}
47
48#[derive(Debug, Clone, Serialize, Deserialize)]
50pub enum DaemonResponse {
51 Initialized,
52 Pong,
53 LspResult { client_id: i64, result: Result<Value, LspErrorResponse> },
54 Error(ProtocolError),
55}
56
57#[derive(Debug, Clone, Serialize, Deserialize)]
59pub struct LspErrorResponse {
60 pub code: i32,
61 pub message: String,
62}
63
64#[derive(Debug, Clone, Serialize, Deserialize)]
66pub struct ProtocolError {
67 pub message: String,
68 #[serde(skip_serializing_if = "Option::is_none")]
70 pub client_id: Option<i64>,
71}
72
73impl ProtocolError {
74 pub fn new(message: impl Into<String>) -> Self {
75 Self { message: message.into(), client_id: None }
76 }
77
78 pub fn with_client_id(message: impl Into<String>, client_id: i64) -> Self {
79 Self { message: message.into(), client_id: Some(client_id) }
80 }
81}
82
83pub fn extract_document_uri(method: &str, params: &Value) -> Option<Uri> {
88 if !method.starts_with("textDocument/") {
89 return None;
90 }
91 params.pointer("/textDocument/uri").and_then(|v| v.as_str()).and_then(|s| s.parse().ok())
92}
93
94pub const MAX_MESSAGE_SIZE: u32 = 16 * 1024 * 1024;
96
97pub const LSP_REQUEST_TIMED_OUT: i32 = -33001;
104
105pub const LSP_TRANSPORT_CLOSED: i32 = -33002;
108
109fn invalid_data(err: impl Into<Box<dyn std::error::Error + Send + Sync>>) -> io::Error {
110 io::Error::new(io::ErrorKind::InvalidData, err)
111}
112
113pub(crate) struct JsonFrames<T>(LengthDelimitedCodec, PhantomData<fn() -> T>);
114
115impl<T> JsonFrames<T> {
116 pub(crate) fn new() -> Self {
117 Self(
118 LengthDelimitedCodec::builder()
119 .big_endian()
120 .length_field_type::<u32>()
121 .max_frame_length(MAX_MESSAGE_SIZE as usize)
122 .new_codec(),
123 PhantomData,
124 )
125 }
126}
127
128impl<T: DeserializeOwned> Decoder for JsonFrames<T> {
129 type Item = T;
130 type Error = io::Error;
131
132 fn decode(&mut self, src: &mut BytesMut) -> io::Result<Option<T>> {
133 self.0.decode(src)?.map(|b| serde_json::from_slice(&b).map_err(invalid_data)).transpose()
134 }
135}
136
137impl<T: Serialize> Encoder<T> for JsonFrames<T> {
138 type Error = io::Error;
139
140 fn encode(&mut self, item: T, dst: &mut BytesMut) -> io::Result<()> {
141 let json = serde_json::to_vec(&item).map_err(invalid_data)?;
142 self.0.encode(Bytes::from(json), dst)
143 }
144}
145
146pub(crate) type FrameReader<R, T> = FramedRead<R, JsonFrames<T>>;
147pub(crate) type FrameWriter<W, T> = FramedWrite<W, JsonFrames<T>>;
148
149pub(crate) fn frame_reader<R: tokio::io::AsyncRead, T: DeserializeOwned>(reader: R) -> FrameReader<R, T> {
150 FramedRead::new(reader, JsonFrames::new())
151}
152
153pub(crate) fn frame_writer<W: tokio::io::AsyncWrite, T: Serialize>(writer: W) -> FrameWriter<W, T> {
154 FramedWrite::new(writer, JsonFrames::new())
155}
156
157#[cfg(test)]
158mod tests {
159 use super::*;
160
161 #[test]
162 fn test_protocol_error_new() {
163 let err = ProtocolError::new("test error");
164 assert_eq!(err.message, "test error");
165 }
166
167 #[test]
168 fn test_daemon_request_lsp_call_roundtrip() {
169 let req = DaemonRequest::LspCall {
170 client_id: 42,
171 method: "textDocument/definition".to_string(),
172 params: serde_json::json!({
173 "textDocument": { "uri": "file:///test.rs" },
174 "position": { "line": 0, "character": 0 }
175 }),
176 };
177 let json = serde_json::to_string(&req).unwrap();
178 let decoded: DaemonRequest = serde_json::from_str(&json).unwrap();
179 match decoded {
180 DaemonRequest::LspCall { client_id, method, .. } => {
181 assert_eq!(client_id, 42);
182 assert_eq!(method, "textDocument/definition");
183 }
184 _ => panic!("Wrong variant"),
185 }
186 }
187
188 #[test]
189 fn test_extract_document_uri_definition() {
190 let params = serde_json::json!({
191 "textDocument": { "uri": "file:///src/main.rs" },
192 "position": { "line": 10, "character": 5 }
193 });
194 let uri = extract_document_uri("textDocument/definition", ¶ms);
195 assert!(uri.is_some());
196 assert_eq!(uri.unwrap().as_str(), "file:///src/main.rs");
197 }
198
199 #[test]
200 fn test_extract_document_uri_references() {
201 let params = serde_json::json!({
202 "textDocument": { "uri": "file:///src/lib.rs" },
203 "position": { "line": 5, "character": 3 },
204 "context": { "includeDeclaration": true }
205 });
206 let uri = extract_document_uri("textDocument/references", ¶ms);
207 assert!(uri.is_some());
208 assert_eq!(uri.unwrap().as_str(), "file:///src/lib.rs");
209 }
210
211 #[test]
212 fn test_extract_document_uri_document_symbol() {
213 let params = serde_json::json!({
214 "textDocument": { "uri": "file:///src/foo.rs" }
215 });
216 let uri = extract_document_uri("textDocument/documentSymbol", ¶ms);
217 assert!(uri.is_some());
218 assert_eq!(uri.unwrap().as_str(), "file:///src/foo.rs");
219 }
220
221 #[test]
222 fn test_extract_document_uri_workspace_symbol_returns_none() {
223 let params = serde_json::json!({ "query": "Foo" });
224 let uri = extract_document_uri("workspace/symbol", ¶ms);
225 assert!(uri.is_none());
226 }
227
228 #[test]
229 fn test_extract_document_uri_unknown_method_returns_none() {
230 let params = serde_json::json!({});
231 let uri = extract_document_uri("textDocument/unknown", ¶ms);
232 assert!(uri.is_none());
233 }
234
235 #[tokio::test]
236 async fn test_duplex_reads_multiple_back_to_back_frames() {
237 use futures::{SinkExt, StreamExt};
238
239 let (client_io, server_io) = tokio::io::duplex(1024);
240 let mut writer = frame_writer::<_, DaemonRequest>(client_io);
241 let mut reader = frame_reader::<_, DaemonRequest>(server_io);
242
243 writer.send(DaemonRequest::Ping).await.expect("send ping");
244 writer.send(DaemonRequest::Disconnect).await.expect("send disconnect");
245
246 let first = reader.next().await.expect("first frame").expect("decode first frame");
247 assert!(matches!(first, DaemonRequest::Ping));
248
249 let second = reader.next().await.expect("second frame").expect("decode second frame");
250 assert!(matches!(second, DaemonRequest::Disconnect));
251 }
252}