Skip to main content

aether_lspd/
protocol.rs

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::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        /// If None, return all cached diagnostics for the workspace
24        uri: Option<Uri>,
25    },
26    QueueDiagnosticRefresh {
27        client_id: i64,
28        uri: Uri,
29    },
30    Disconnect,
31    Ping,
32}
33
34/// Initialize request to set up LSP for a workspace
35#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct InitializeRequest {
37    pub workspace_root: PathBuf,
38    pub language: LanguageId,
39}
40
41/// LSP notification from client to server
42#[derive(Debug, Clone, Serialize, Deserialize)]
43pub struct LspNotification {
44    pub method: String,
45    pub params: Value,
46}
47
48/// Top-level daemon response
49#[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/// LSP error response
58#[derive(Debug, Clone, Serialize, Deserialize)]
59pub struct LspErrorResponse {
60    pub code: i32,
61    pub message: String,
62}
63
64/// Protocol-level error (not LSP error)
65#[derive(Debug, Clone, Serialize, Deserialize)]
66pub struct ProtocolError {
67    pub message: String,
68    /// Optional `client_id` for correlating errors back to LSP requests
69    #[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
83/// Extract the document URI from an LSP request's params by method name.
84///
85/// Used by the daemon for auto-open: if the request targets a specific file,
86/// the daemon ensures the file is opened before forwarding the request.
87pub 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
94/// Maximum message size (16 MB)
95pub const MAX_MESSAGE_SIZE: u32 = 16 * 1024 * 1024;
96
97/// Error code for a request that exceeded the daemon's per-request timeout.
98/// The daemon kills and replaces the language server when this fires.
99///
100/// Daemon-synthesized codes live outside the ranges reserved by JSON-RPC
101/// (-32000..=-32768) and LSP (-32800..=-32899), so they can never collide with
102/// a code forwarded verbatim from a real language server.
103pub const LSP_REQUEST_TIMED_OUT: i32 = -33001;
104
105/// Error code for a request that failed because the language server process
106/// exited or its transport shut down before responding.
107pub 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(json.into(), 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", &params);
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", &params);
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", &params);
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", &params);
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", &params);
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}