Skip to main content

shuck_server/
server.rs

1use lsp_server::Connection;
2use lsp_types as types;
3use lsp_types::InitializeParams;
4use lsp_types::{
5    ClientCapabilities, CodeActionKind, CodeActionOptions, DiagnosticOptions, OneOf,
6    TextDocumentSyncCapability, TextDocumentSyncKind, TextDocumentSyncOptions,
7    WorkDoneProgressOptions, WorkspaceFoldersServerCapabilities,
8};
9use std::num::NonZeroUsize;
10
11pub(crate) use self::connection::ConnectionInitializer;
12pub use self::connection::ConnectionSender;
13use self::schedule::spawn_main_loop;
14use crate::PositionEncoding;
15pub use crate::server::main_loop::MainLoopSender;
16pub(crate) use crate::server::main_loop::{Event, MainLoopReceiver};
17use crate::session::{AllOptions, Client, Session};
18use crate::workspace::Workspaces;
19pub(crate) use api::Error;
20
21mod api;
22mod connection;
23mod main_loop;
24mod schedule;
25
26pub(crate) type Result<T> = std::result::Result<T, api::Error>;
27
28/// Initialized Shuck LSP server.
29pub struct Server {
30    connection: Connection,
31    client_capabilities: ClientCapabilities,
32    worker_threads: NonZeroUsize,
33    main_loop_receiver: MainLoopReceiver,
34    main_loop_sender: MainLoopSender,
35    session: Session,
36}
37
38impl Server {
39    pub(crate) fn new(
40        worker_threads: NonZeroUsize,
41        connection: ConnectionInitializer,
42    ) -> crate::Result<Self> {
43        let (id, init_params) = connection.initialize_start()?;
44        let client_capabilities = init_params.capabilities;
45        let position_encoding = Self::find_best_position_encoding(&client_capabilities);
46        #[allow(deprecated)]
47        let InitializeParams {
48            initialization_options,
49            root_path,
50            root_uri,
51            workspace_folders,
52            ..
53        } = init_params;
54        let all_options =
55            AllOptions::from_value(initialization_options.unwrap_or(serde_json::Value::Null));
56        let workspace_diagnostics_enabled = all_options.workspace_diagnostics_enabled();
57        let AllOptions { global, workspace } = all_options;
58        let server_capabilities =
59            Self::server_capabilities(position_encoding, workspace_diagnostics_enabled);
60        let connection = connection.initialize_finish(
61            id,
62            &server_capabilities,
63            crate::SERVER_NAME,
64            crate::version(),
65        )?;
66
67        let (main_loop_sender, main_loop_receiver) = crossbeam::channel::bounded(32);
68
69        let client = Client::new(main_loop_sender.clone(), connection.sender.clone());
70
71        crate::logging::init_logging(
72            global.tracing.log_level.unwrap_or_default(),
73            global.tracing.log_file.as_deref(),
74        );
75
76        let workspaces = Workspaces::from_workspace_folders(
77            workspace_folders,
78            root_uri,
79            root_path,
80            workspace.unwrap_or_default(),
81        )?;
82        let global = global.into_settings(client.clone());
83
84        Ok(Self {
85            connection,
86            client_capabilities: client_capabilities.clone(),
87            worker_threads,
88            main_loop_receiver,
89            main_loop_sender,
90            session: Session::new(
91                &client_capabilities,
92                position_encoding,
93                global,
94                &workspaces,
95                &client,
96            )?,
97        })
98    }
99
100    /// Run the server main loop until shutdown or error.
101    pub fn run(mut self) -> crate::Result<()> {
102        let panic_client = Client::new(
103            self.main_loop_sender.clone(),
104            self.connection.sender.clone(),
105        );
106        let _panic_hook = install_panic_hook(panic_client);
107        spawn_main_loop(move || self.main_loop())?
108            .join()
109            .map_err(|_| anyhow::anyhow!("main loop thread panicked"))?
110    }
111
112    fn find_best_position_encoding(client_capabilities: &ClientCapabilities) -> PositionEncoding {
113        client_capabilities
114            .general
115            .as_ref()
116            .and_then(|general| general.position_encodings.as_ref())
117            .and_then(|encodings| {
118                encodings
119                    .iter()
120                    .filter_map(|encoding| PositionEncoding::try_from(encoding).ok())
121                    .max()
122            })
123            .unwrap_or_default()
124    }
125
126    fn server_capabilities(
127        position_encoding: PositionEncoding,
128        workspace_diagnostics_enabled: bool,
129    ) -> types::ServerCapabilities {
130        types::ServerCapabilities {
131            position_encoding: Some(position_encoding.into()),
132            code_action_provider: Some(types::CodeActionProviderCapability::Options(
133                CodeActionOptions {
134                    code_action_kinds: Some(
135                        SupportedCodeAction::all()
136                            .map(SupportedCodeAction::to_kind)
137                            .collect(),
138                    ),
139                    work_done_progress_options: WorkDoneProgressOptions {
140                        work_done_progress: Some(true),
141                    },
142                    resolve_provider: Some(true),
143                },
144            )),
145            workspace: Some(types::WorkspaceServerCapabilities {
146                workspace_folders: Some(WorkspaceFoldersServerCapabilities {
147                    supported: Some(true),
148                    change_notifications: Some(OneOf::Left(true)),
149                }),
150                file_operations: None,
151            }),
152            completion_provider: Some(types::CompletionOptions {
153                resolve_provider: Some(true),
154                trigger_characters: Some(vec!["$".to_owned(), "{".to_owned()]),
155                ..types::CompletionOptions::default()
156            }),
157            definition_provider: Some(OneOf::Left(true)),
158            document_link_provider: Some(types::DocumentLinkOptions {
159                resolve_provider: Some(false),
160                work_done_progress_options: WorkDoneProgressOptions {
161                    work_done_progress: Some(true),
162                },
163            }),
164            call_hierarchy_provider: Some(types::CallHierarchyServerCapability::Simple(true)),
165            references_provider: Some(OneOf::Left(true)),
166            document_highlight_provider: Some(OneOf::Left(true)),
167            document_formatting_provider: Some(OneOf::Left(true)),
168            document_range_formatting_provider: Some(OneOf::Left(true)),
169            folding_range_provider: Some(types::FoldingRangeProviderCapability::Simple(true)),
170            document_symbol_provider: Some(OneOf::Left(true)),
171            workspace_symbol_provider: Some(OneOf::Right(types::WorkspaceSymbolOptions {
172                work_done_progress_options: WorkDoneProgressOptions {
173                    work_done_progress: Some(true),
174                },
175                resolve_provider: Some(false),
176            })),
177            diagnostic_provider: Some(types::DiagnosticServerCapabilities::Options(
178                DiagnosticOptions {
179                    identifier: Some(crate::DIAGNOSTIC_NAME.into()),
180                    inter_file_dependencies: false,
181                    workspace_diagnostics: workspace_diagnostics_enabled,
182                    work_done_progress_options: WorkDoneProgressOptions {
183                        work_done_progress: Some(true),
184                    },
185                },
186            )),
187            execute_command_provider: Some(types::ExecuteCommandOptions {
188                commands: SupportedCommand::all()
189                    .map(|command| command.identifier().to_string())
190                    .collect(),
191                work_done_progress_options: WorkDoneProgressOptions {
192                    work_done_progress: Some(false),
193                },
194            }),
195            hover_provider: Some(types::HoverProviderCapability::Simple(true)),
196            rename_provider: Some(OneOf::Right(types::RenameOptions {
197                prepare_provider: Some(true),
198                work_done_progress_options: WorkDoneProgressOptions {
199                    work_done_progress: Some(true),
200                },
201            })),
202            selection_range_provider: Some(types::SelectionRangeProviderCapability::Simple(true)),
203            text_document_sync: Some(TextDocumentSyncCapability::Options(
204                TextDocumentSyncOptions {
205                    open_close: Some(true),
206                    change: Some(TextDocumentSyncKind::INCREMENTAL),
207                    will_save: Some(false),
208                    will_save_wait_until: Some(false),
209                    ..Default::default()
210                },
211            )),
212            ..Default::default()
213        }
214    }
215}
216
217#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
218pub(crate) enum SupportedCodeAction {
219    QuickFix,
220    SourceFixAll,
221}
222
223impl SupportedCodeAction {
224    fn all() -> impl Iterator<Item = Self> {
225        [Self::QuickFix, Self::SourceFixAll].into_iter()
226    }
227
228    fn to_kind(self) -> CodeActionKind {
229        match self {
230            Self::QuickFix => CodeActionKind::QUICKFIX,
231            Self::SourceFixAll => crate::SOURCE_FIX_ALL_SHUCK,
232        }
233    }
234}
235
236#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
237pub(crate) enum SupportedCommand {
238    ApplyAutofix,
239    ApplyDirective,
240    PrintDebugInformation,
241}
242
243impl SupportedCommand {
244    fn all() -> impl Iterator<Item = Self> {
245        [
246            Self::ApplyAutofix,
247            Self::ApplyDirective,
248            Self::PrintDebugInformation,
249        ]
250        .into_iter()
251    }
252
253    fn identifier(self) -> &'static str {
254        match self {
255            Self::ApplyAutofix => "shuck.applyAutofix",
256            Self::ApplyDirective => "shuck.applyDirective",
257            Self::PrintDebugInformation => "shuck.printDebugInformation",
258        }
259    }
260}
261
262type PanicHook = Box<dyn Fn(&std::panic::PanicHookInfo<'_>) + Sync + Send + 'static>;
263
264struct PanicHookGuard {
265    previous: Option<PanicHook>,
266}
267
268impl Drop for PanicHookGuard {
269    fn drop(&mut self) {
270        if let Some(previous) = self.previous.take() {
271            std::panic::set_hook(previous);
272        }
273    }
274}
275
276fn install_panic_hook(client: Client) -> PanicHookGuard {
277    let previous = std::panic::take_hook();
278    std::panic::set_hook(Box::new(move |panic_info| {
279        report_panic(&client, panic_info);
280    }));
281    PanicHookGuard {
282        previous: Some(previous),
283    }
284}
285
286fn report_panic(client: &Client, panic_info: &std::panic::PanicHookInfo<'_>) {
287    let summary = panic_info
288        .payload()
289        .downcast_ref::<String>()
290        .cloned()
291        .or_else(|| {
292            panic_info
293                .payload()
294                .downcast_ref::<&'static str>()
295                .map(|message| (*message).to_owned())
296        })
297        .unwrap_or_else(|| "unknown panic".to_owned());
298    let location = panic_info.location().map(|location| {
299        format!(
300            "{}:{}:{}",
301            location.file(),
302            location.line(),
303            location.column()
304        )
305    });
306    let backtrace = std::backtrace::Backtrace::force_capture().to_string();
307    emit_panic_report(client, &summary, location.as_deref(), &backtrace);
308}
309
310fn emit_panic_report(client: &Client, summary: &str, location: Option<&str>, backtrace: &str) {
311    let location = location.unwrap_or("unknown location");
312    let details = format!("Shuck server panicked at {location}: {summary}\n{backtrace}");
313    tracing::error!("{details}");
314    eprintln!("{details}");
315    if let Err(error) = client.log_message(&details, lsp_types::MessageType::ERROR) {
316        tracing::error!("Failed to send panic log message to client: {error}");
317    }
318    client.show_error_message(format!("Shuck server panicked: {summary}"));
319}
320
321#[cfg(test)]
322mod tests {
323    use crossbeam::channel;
324    use lsp_server::Message;
325    use lsp_types::notification::Notification;
326
327    use super::*;
328    use crate::Client;
329
330    #[test]
331    fn advertises_formatting_capabilities() {
332        let capabilities = Server::server_capabilities(PositionEncoding::UTF16, false);
333        assert_eq!(
334            capabilities.document_formatting_provider,
335            Some(OneOf::Left(true))
336        );
337        assert_eq!(
338            capabilities.document_range_formatting_provider,
339            Some(OneOf::Left(true))
340        );
341    }
342
343    #[test]
344    fn advertises_navigation_completion_and_rename_capabilities() {
345        let capabilities = Server::server_capabilities(PositionEncoding::UTF16, false);
346        assert!(capabilities.completion_provider.is_some());
347        assert_eq!(capabilities.definition_provider, Some(OneOf::Left(true)));
348        assert!(capabilities.document_link_provider.is_some());
349        assert_eq!(capabilities.references_provider, Some(OneOf::Left(true)));
350        assert_eq!(
351            capabilities.document_highlight_provider,
352            Some(OneOf::Left(true))
353        );
354        let Some(OneOf::Right(rename)) = capabilities.rename_provider else {
355            panic!("expected rename options");
356        };
357        assert_eq!(rename.prepare_provider, Some(true));
358    }
359
360    #[test]
361    fn advertises_document_symbol_capability() {
362        let capabilities = Server::server_capabilities(PositionEncoding::UTF16, false);
363        assert_eq!(
364            capabilities.document_symbol_provider,
365            Some(OneOf::Left(true))
366        );
367    }
368
369    #[test]
370    fn advertises_workspace_symbol_capability_without_resolve() {
371        let capabilities = Server::server_capabilities(PositionEncoding::UTF16, false);
372        let Some(OneOf::Right(options)) = capabilities.workspace_symbol_provider else {
373            panic!("expected workspace symbol options");
374        };
375        assert_eq!(options.resolve_provider, Some(false));
376    }
377
378    #[test]
379    fn advertises_only_non_formatting_execute_commands() {
380        let capabilities = Server::server_capabilities(PositionEncoding::UTF16, false);
381        let commands = capabilities
382            .execute_command_provider
383            .expect("server should advertise execute commands")
384            .commands;
385
386        assert!(commands.contains(&"shuck.applyAutofix".to_owned()));
387        assert!(commands.contains(&"shuck.applyDirective".to_owned()));
388        assert!(commands.contains(&"shuck.printDebugInformation".to_owned()));
389        assert!(!commands.contains(&"shuck.applyFormat".to_owned()));
390    }
391
392    #[test]
393    fn advertises_workspace_diagnostics_only_when_enabled() {
394        for (enabled, expected) in [(false, false), (true, true)] {
395            let capabilities = Server::server_capabilities(PositionEncoding::UTF16, enabled);
396            let Some(types::DiagnosticServerCapabilities::Options(options)) =
397                capabilities.diagnostic_provider
398            else {
399                panic!("expected diagnostic options");
400            };
401            assert_eq!(options.workspace_diagnostics, expected);
402        }
403    }
404
405    #[test]
406    fn panic_reports_are_sent_to_the_client() {
407        let (main_loop_sender, _main_loop_receiver) = channel::unbounded();
408        let (client_sender, client_receiver) = channel::unbounded();
409        let client = Client::new(main_loop_sender, client_sender);
410
411        emit_panic_report(&client, "boom", Some("test.rs:1:1"), "stack backtrace");
412
413        let first = client_receiver
414            .recv_timeout(std::time::Duration::from_secs(1))
415            .expect("panic log notification should be sent");
416        let second = client_receiver
417            .recv_timeout(std::time::Duration::from_secs(1))
418            .expect("panic showMessage notification should be sent");
419
420        let messages = [first, second];
421        assert!(messages.iter().any(|message| matches!(
422            message,
423            Message::Notification(notification)
424                if notification.method == lsp_types::notification::LogMessage::METHOD
425        )));
426        assert!(messages.iter().any(|message| matches!(
427            message,
428            Message::Notification(notification)
429                if notification.method == lsp_types::notification::ShowMessage::METHOD
430        )));
431    }
432}