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