Skip to main content

aether_lspd/
client.rs

1use std::collections::HashMap;
2use std::io;
3use std::io::ErrorKind;
4use std::path::{Path, PathBuf};
5use std::process::Stdio;
6use std::sync::atomic::{AtomicI64, Ordering};
7use std::sync::{Arc, Mutex, PoisonError};
8use std::time::Duration;
9
10use futures::{SinkExt, StreamExt};
11use lsp_types::{
12    CallHierarchyIncomingCall, CallHierarchyIncomingCallsParams, CallHierarchyItem, CallHierarchyOutgoingCall,
13    CallHierarchyOutgoingCallsParams, CallHierarchyPrepareParams, DocumentSymbolParams, DocumentSymbolResponse,
14    GotoDefinitionParams, GotoDefinitionResponse, Hover, HoverParams, Location, PartialResultParams, Position,
15    PublishDiagnosticsParams, ReferenceContext, ReferenceParams, RenameParams, SymbolInformation,
16    TextDocumentIdentifier, TextDocumentPositionParams, Uri, WorkDoneProgressParams, WorkspaceEdit,
17    WorkspaceSymbolParams,
18};
19use serde::Serialize;
20use serde::de::DeserializeOwned;
21use serde_json::Value;
22use thiserror::Error;
23use tokio::io::{ReadHalf, WriteHalf};
24use tokio::net::UnixStream;
25use tokio::process::Command;
26use tokio::sync::{Mutex as AsyncMutex, oneshot};
27
28use crate::language_catalog::LanguageId;
29use crate::protocol::{DaemonRequest, DaemonResponse, FrameWriter, InitializeRequest, frame_reader, frame_writer};
30use crate::socket_path::{ensure_socket_dir, log_file_path};
31
32#[doc = include_str!("docs/client_error.md")]
33#[derive(Debug, Error)]
34pub enum ClientError {
35    #[error("Failed to connect to daemon: {0}")]
36    ConnectionFailed(#[source] io::Error),
37
38    #[error("IO error: {0}")]
39    Io(#[from] io::Error),
40
41    #[error("Daemon error: {0}")]
42    DaemonError(String),
43
44    #[error("LSP error (code={code}): {message}")]
45    LspError { code: i32, message: String },
46
47    #[error("Failed to spawn daemon: {0}")]
48    SpawnFailed(#[source] io::Error),
49
50    #[error("Timeout waiting for daemon to start")]
51    SpawnTimeout,
52
53    #[error("Daemon binary not found: {0}")]
54    DaemonBinaryNotFound(String),
55
56    #[error("Protocol error: {0}")]
57    ProtocolError(String),
58
59    #[error("Initialization failed: {0}")]
60    InitializationFailed(String),
61}
62
63pub type ClientResult<T> = std::result::Result<T, ClientError>;
64
65#[doc = include_str!("docs/client.md")]
66pub struct LspClient {
67    writer: AsyncMutex<FrameWriter<WriteHalf<UnixStream>, DaemonRequest>>,
68    pending: PendingRequests,
69    next_id: AtomicI64,
70    reader_task: tokio::task::JoinHandle<()>,
71}
72
73impl LspClient {
74    pub async fn connect(workspace_root: &Path, language: LanguageId) -> ClientResult<Self> {
75        let socket_path = ensure_socket_dir(workspace_root, language).map_err(ClientError::Io)?;
76
77        match UnixStream::connect(&socket_path).await {
78            Ok(stream) => {
79                return Self::from_stream(stream, workspace_root, language).await;
80            }
81            Err(err) if err.kind() == ErrorKind::ConnectionRefused || err.kind() == ErrorKind::NotFound => {}
82            Err(err) => return Err(ClientError::ConnectionFailed(err)),
83        }
84
85        spawn_daemon(&socket_path).await?;
86        let stream = UnixStream::connect(&socket_path).await.map_err(ClientError::ConnectionFailed)?;
87        Self::from_stream(stream, workspace_root, language).await
88    }
89
90    pub async fn goto_definition(&self, uri: Uri, line: u32, character: u32) -> ClientResult<GotoDefinitionResponse> {
91        let params = GotoDefinitionParams {
92            text_document_position_params: TextDocumentPositionParams {
93                text_document: TextDocumentIdentifier { uri },
94                position: Position { line, character },
95            },
96            work_done_progress_params: WorkDoneProgressParams::default(),
97            partial_result_params: PartialResultParams::default(),
98        };
99        self.call("textDocument/definition", &params, || GotoDefinitionResponse::Array(vec![])).await
100    }
101
102    pub async fn goto_implementation(
103        &self,
104        uri: Uri,
105        line: u32,
106        character: u32,
107    ) -> ClientResult<GotoDefinitionResponse> {
108        let params = GotoDefinitionParams {
109            text_document_position_params: TextDocumentPositionParams {
110                text_document: TextDocumentIdentifier { uri },
111                position: Position { line, character },
112            },
113            work_done_progress_params: WorkDoneProgressParams::default(),
114            partial_result_params: PartialResultParams::default(),
115        };
116        self.call("textDocument/implementation", &params, || GotoDefinitionResponse::Array(vec![])).await
117    }
118
119    pub async fn find_references(
120        &self,
121        uri: Uri,
122        line: u32,
123        character: u32,
124        include_declaration: bool,
125    ) -> ClientResult<Vec<Location>> {
126        let params = ReferenceParams {
127            text_document_position: TextDocumentPositionParams {
128                text_document: TextDocumentIdentifier { uri },
129                position: Position { line, character },
130            },
131            work_done_progress_params: WorkDoneProgressParams::default(),
132            partial_result_params: PartialResultParams::default(),
133            context: ReferenceContext { include_declaration },
134        };
135        self.call("textDocument/references", &params, Vec::new).await
136    }
137
138    pub async fn hover(&self, uri: Uri, line: u32, character: u32) -> ClientResult<Option<Hover>> {
139        let params = HoverParams {
140            text_document_position_params: TextDocumentPositionParams {
141                text_document: TextDocumentIdentifier { uri },
142                position: Position { line, character },
143            },
144            work_done_progress_params: WorkDoneProgressParams::default(),
145        };
146        self.call("textDocument/hover", &params, || None).await
147    }
148
149    pub async fn workspace_symbol(&self, query: String) -> ClientResult<Vec<SymbolInformation>> {
150        let params = WorkspaceSymbolParams {
151            query,
152            partial_result_params: PartialResultParams::default(),
153            work_done_progress_params: WorkDoneProgressParams::default(),
154        };
155        self.call("workspace/symbol", &params, Vec::new).await
156    }
157
158    pub async fn document_symbol(&self, uri: Uri) -> ClientResult<DocumentSymbolResponse> {
159        let params = DocumentSymbolParams {
160            text_document: TextDocumentIdentifier { uri },
161            work_done_progress_params: WorkDoneProgressParams::default(),
162            partial_result_params: PartialResultParams::default(),
163        };
164        self.call("textDocument/documentSymbol", &params, || DocumentSymbolResponse::Flat(vec![])).await
165    }
166
167    pub async fn prepare_call_hierarchy(
168        &self,
169        uri: Uri,
170        line: u32,
171        character: u32,
172    ) -> ClientResult<Vec<CallHierarchyItem>> {
173        let params = CallHierarchyPrepareParams {
174            text_document_position_params: TextDocumentPositionParams {
175                text_document: TextDocumentIdentifier { uri },
176                position: Position { line, character },
177            },
178            work_done_progress_params: WorkDoneProgressParams::default(),
179        };
180        self.call("textDocument/prepareCallHierarchy", &params, Vec::new).await
181    }
182
183    pub async fn incoming_calls(&self, item: CallHierarchyItem) -> ClientResult<Vec<CallHierarchyIncomingCall>> {
184        let params = CallHierarchyIncomingCallsParams {
185            item,
186            work_done_progress_params: WorkDoneProgressParams::default(),
187            partial_result_params: PartialResultParams::default(),
188        };
189        self.call("callHierarchy/incomingCalls", &params, Vec::new).await
190    }
191
192    pub async fn outgoing_calls(&self, item: CallHierarchyItem) -> ClientResult<Vec<CallHierarchyOutgoingCall>> {
193        let params = CallHierarchyOutgoingCallsParams {
194            item,
195            work_done_progress_params: WorkDoneProgressParams::default(),
196            partial_result_params: PartialResultParams::default(),
197        };
198        self.call("callHierarchy/outgoingCalls", &params, Vec::new).await
199    }
200
201    pub async fn rename(
202        &self,
203        uri: Uri,
204        line: u32,
205        character: u32,
206        new_name: String,
207    ) -> ClientResult<Option<WorkspaceEdit>> {
208        let params = RenameParams {
209            text_document_position: TextDocumentPositionParams {
210                text_document: TextDocumentIdentifier { uri },
211                position: Position { line, character },
212            },
213            new_name,
214            work_done_progress_params: WorkDoneProgressParams::default(),
215        };
216        self.call("textDocument/rename", &params, || None).await
217    }
218
219    pub async fn get_diagnostics(&self, uri: Option<Uri>) -> ClientResult<Vec<PublishDiagnosticsParams>> {
220        let client_id = self.next_id.fetch_add(1, Ordering::SeqCst);
221        let request = DaemonRequest::GetDiagnostics { client_id, uri };
222
223        self.send_and_await(request, client_id)
224            .await
225            .and_then(|value| serde_json::from_value(value).map_err(|err| ClientError::ProtocolError(err.to_string())))
226    }
227
228    pub async fn queue_diagnostic_refresh(&self, uri: Uri) -> ClientResult<()> {
229        let client_id = self.next_id.fetch_add(1, Ordering::SeqCst);
230        let request = DaemonRequest::QueueDiagnosticRefresh { client_id, uri };
231        self.send_and_await(request, client_id).await.map(|_| ())
232    }
233
234    /// Whether the daemon connection's reader task is still running.
235    pub fn is_connected(&self) -> bool {
236        !self.reader_task.is_finished()
237    }
238
239    pub async fn disconnect(self) -> ClientResult<()> {
240        let mut writer = self.writer.lock().await;
241        writer.send(DaemonRequest::Disconnect).await.map_err(ClientError::Io)
242    }
243
244    pub async fn call<P: Serialize, R: DeserializeOwned>(
245        &self,
246        method: &str,
247        params: &P,
248        default: impl FnOnce() -> R,
249    ) -> ClientResult<R> {
250        let params_value = serde_json::to_value(params).map_err(|err| ClientError::ProtocolError(err.to_string()))?;
251
252        let client_id = self.next_id.fetch_add(1, Ordering::SeqCst);
253        let request = DaemonRequest::LspCall { client_id, method: method.to_string(), params: params_value };
254
255        let value = self.send_and_await(request, client_id).await?;
256
257        if value.is_null() {
258            Ok(default())
259        } else {
260            serde_json::from_value(value).map_err(|err| ClientError::ProtocolError(format!("Parse error: {err}")))
261        }
262    }
263}
264
265impl LspClient {
266    async fn from_stream(stream: UnixStream, workspace_root: &Path, language: LanguageId) -> ClientResult<Self> {
267        let (reader, writer) = tokio::io::split(stream);
268        let mut reader = frame_reader::<_, DaemonResponse>(reader);
269        let mut writer = frame_writer::<_, DaemonRequest>(writer);
270
271        let initialize =
272            DaemonRequest::Initialize(InitializeRequest { workspace_root: workspace_root.to_path_buf(), language });
273
274        writer.send(initialize).await.map_err(ClientError::Io)?;
275
276        let response = match reader.next().await {
277            Some(Ok(resp)) => resp,
278            Some(Err(err)) => return Err(ClientError::Io(err)),
279            None => {
280                return Err(ClientError::ProtocolError("Connection closed during initialization".into()));
281            }
282        };
283
284        match response {
285            DaemonResponse::Initialized => {}
286            DaemonResponse::Error(err) => {
287                return Err(ClientError::InitializationFailed(err.message));
288            }
289            _ => {
290                return Err(ClientError::ProtocolError("Unexpected response to Initialize".into()));
291            }
292        }
293
294        let pending = Arc::new(Mutex::new(Some(HashMap::new())));
295        let reader_task = tokio::spawn(run_reader(reader, Arc::clone(&pending)));
296
297        Ok(Self { writer: AsyncMutex::new(writer), pending, next_id: AtomicI64::new(1), reader_task })
298    }
299
300    async fn send_and_await(&self, request: DaemonRequest, client_id: i64) -> ClientResult<Value> {
301        let (response_tx, response_rx) = oneshot::channel();
302        {
303            let mut pending = self.pending.lock().unwrap_or_else(PoisonError::into_inner);
304            let pending = pending.as_mut().ok_or_else(|| ClientError::ProtocolError("Daemon disconnected".into()))?;
305            pending.insert(client_id, response_tx);
306        }
307
308        let write_result = {
309            let mut writer = self.writer.lock().await;
310            writer.send(request).await
311        };
312        if let Err(error) = write_result {
313            if let Some(pending) = self.pending.lock().unwrap_or_else(PoisonError::into_inner).as_mut() {
314                pending.remove(&client_id);
315            }
316            return Err(ClientError::Io(error));
317        }
318        response_rx.await.map_err(|_| ClientError::ProtocolError("Response channel closed".into()))?
319    }
320}
321
322impl Drop for LspClient {
323    fn drop(&mut self) {
324        self.reader_task.abort();
325    }
326}
327
328type PendingResult = Result<Value, ClientError>;
329type PendingRequests = Arc<Mutex<Option<HashMap<i64, oneshot::Sender<PendingResult>>>>>;
330
331async fn run_reader(
332    mut reader: crate::protocol::FrameReader<ReadHalf<UnixStream>, DaemonResponse>,
333    pending: PendingRequests,
334) {
335    while let Some(msg) = reader.next().await {
336        let response = match msg {
337            Ok(response) => response,
338            Err(error) => {
339                tracing::debug!(%error, "Error reading daemon response");
340                break;
341            }
342        };
343        let (client_id, result) = match response {
344            DaemonResponse::LspResult { client_id, result } => {
345                (client_id, result.map_err(|error| ClientError::LspError { code: error.code, message: error.message }))
346            }
347            DaemonResponse::Error(error) => {
348                let Some(client_id) = error.client_id else { continue };
349                (client_id, Err(ClientError::DaemonError(error.message)))
350            }
351            _ => continue,
352        };
353        if let Some(response_tx) = pending
354            .lock()
355            .unwrap_or_else(PoisonError::into_inner)
356            .as_mut()
357            .and_then(|pending| pending.remove(&client_id))
358        {
359            let _ = response_tx.send(result);
360        }
361    }
362
363    if let Some(pending) = pending.lock().unwrap_or_else(PoisonError::into_inner).take() {
364        for (_, response_tx) in pending {
365            let _ = response_tx.send(Err(ClientError::ProtocolError("Daemon disconnected".into())));
366        }
367    }
368}
369
370async fn spawn_daemon(socket_path: &Path) -> ClientResult<()> {
371    let (binary, subcommand) = find_daemon_binary()?;
372    let log_file = log_file_path(socket_path);
373
374    let mut cmd = Command::new(&binary);
375    if let Some(sub) = subcommand {
376        cmd.arg(sub);
377    }
378    cmd.arg("--socket")
379        .arg(socket_path)
380        .arg("--log-file")
381        .arg(&log_file)
382        .arg("--log-level")
383        .arg("debug")
384        .stdin(Stdio::null())
385        .stdout(Stdio::null())
386        .stderr(Stdio::null());
387
388    #[cfg(unix)]
389    unsafe {
390        use std::os::unix::process::CommandExt;
391        cmd.as_std_mut()
392            .pre_exec(|| nix::unistd::setsid().map(|_| ()).map_err(|e| std::io::Error::from_raw_os_error(e as i32)));
393    }
394
395    let mut child = cmd.spawn().map_err(ClientError::SpawnFailed)?;
396
397    for _ in 0..50 {
398        match child.try_wait() {
399            Ok(Some(status)) if !status.success() => {
400                return Err(ClientError::SpawnFailed(io::Error::other(format!("Daemon exited with status: {status}"))));
401            }
402            Ok(_) => {}
403            Err(err) => return Err(ClientError::SpawnFailed(err)),
404        }
405
406        tokio::time::sleep(Duration::from_millis(100)).await;
407        if UnixStream::connect(socket_path).await.is_ok() {
408            tokio::spawn(async move {
409                match child.wait().await {
410                    Ok(status) => tracing::debug!(%status, "aether-lspd launcher reaped"),
411                    Err(err) => tracing::warn!(%err, "Failed to reap aether-lspd launcher"),
412                }
413            });
414            return Ok(());
415        }
416    }
417
418    let _ = child.kill().await;
419    let _ = child.wait().await;
420    Err(ClientError::SpawnTimeout)
421}
422
423fn find_daemon_binary() -> ClientResult<(PathBuf, Option<&'static str>)> {
424    let exe = std::env::current_exe().ok();
425    let exe_dir = exe.as_deref().and_then(|p| p.parent());
426
427    let standalone_candidates = [
428        exe_dir.map(|dir| dir.join("aether-lspd")),
429        exe_dir.and_then(|dir| dir.parent()).map(|dir| dir.join("aether-lspd")),
430        which_aether_lspd(),
431        Some(PathBuf::from("target/debug/aether-lspd")),
432        Some(PathBuf::from("target/release/aether-lspd")),
433        Some(PathBuf::from("../../target/debug/aether-lspd")),
434        Some(PathBuf::from("../../target/release/aether-lspd")),
435    ];
436
437    for candidate in standalone_candidates.into_iter().flatten() {
438        if candidate.exists() {
439            return Ok((candidate, None));
440        }
441    }
442
443    if let Some(exe) = exe {
444        return Ok((exe, Some("lspd")));
445    }
446
447    Err(ClientError::DaemonBinaryNotFound("aether-lspd not found".into()))
448}
449
450fn which_aether_lspd() -> Option<PathBuf> {
451    std::env::var_os("PATH")
452        .and_then(|paths| std::env::split_paths(&paths).map(|path| path.join("aether-lspd")).find(|path| path.exists()))
453}