Skip to main content

shape_lsp/
trait_lookup.rs

1//! Shared trait lookup helpers for LSP hover/completion.
2//!
3//! Resolves trait definitions from:
4//! 1. Current file AST (fast path)
5//! 2. Importable modules via ModuleCache (stdlib/project modules)
6//!
7//! Also enumerates trait `impl` blocks for a given target type (used by hover
8//! to render an "Implementations for ..." section similar to rust-analyzer).
9
10use crate::doc_render::render_doc_comment;
11use crate::module_cache::ModuleCache;
12use shape_ast::ast::{Item, Program, Span, TraitDef, TypeName};
13use std::path::{Path, PathBuf};
14
15#[derive(Debug, Clone)]
16pub struct ResolvedTraitDef {
17    pub trait_def: TraitDef,
18    pub span: Span,
19    pub documentation: Option<String>,
20    pub source_path: Option<PathBuf>,
21    pub source_text: Option<String>,
22    pub import_path: Option<String>,
23}
24
25pub fn resolve_trait_definition(
26    program: &Program,
27    trait_name: &str,
28    module_cache: Option<&ModuleCache>,
29    current_file: Option<&Path>,
30    workspace_root: Option<&Path>,
31) -> Option<ResolvedTraitDef> {
32    for item in &program.items {
33        if let Item::Trait(trait_def, span) = item {
34            if trait_def.name == trait_name {
35                return Some(ResolvedTraitDef {
36                    trait_def: trait_def.clone(),
37                    span: *span,
38                    documentation: program
39                        .docs
40                        .comment_for_span(*span)
41                        .map(|comment| render_doc_comment(program, comment, None, None, None)),
42                    source_path: None,
43                    source_text: None,
44                    import_path: None,
45                });
46            }
47        }
48    }
49
50    let cache = module_cache?;
51    let current_file = current_file?;
52
53    let mut import_paths = cache.list_importable_modules_with_context(current_file, workspace_root);
54    import_paths.sort();
55
56    for import_path in import_paths {
57        let Some(resolved) = cache.resolve_import(&import_path, current_file, workspace_root)
58        else {
59            continue;
60        };
61        let Some(module_info) =
62            cache.load_module_with_context(&resolved, current_file, workspace_root)
63        else {
64            continue;
65        };
66
67        for item in &module_info.program.items {
68            if let Item::Trait(trait_def, span) = item {
69                if trait_def.name == trait_name {
70                    return Some(ResolvedTraitDef {
71                        trait_def: trait_def.clone(),
72                        span: *span,
73                        documentation: module_info.program.docs.comment_for_span(*span).map(
74                            |comment| {
75                                render_doc_comment(
76                                    &module_info.program,
77                                    comment,
78                                    Some(cache),
79                                    Some(&module_info.path),
80                                    workspace_root,
81                                )
82                            },
83                        ),
84                        source_text: std::fs::read_to_string(&module_info.path).ok(),
85                        source_path: Some(module_info.path.clone()),
86                        import_path: Some(import_path),
87                    });
88                }
89            }
90        }
91    }
92
93    None
94}
95
96/// Summary of a single `impl Trait for Type` block discovered for a target
97/// type. Surfaced by hover to mimic rust-analyzer's "Implementations for ..."
98/// section. `source_module` is `None` when the impl lives in the current file.
99#[derive(Debug, Clone)]
100pub struct ImplSummary {
101    pub trait_name: String,
102    pub impl_name: Option<String>,
103    pub source_module: Option<String>,
104}
105
106fn type_name_base(type_name: &TypeName) -> String {
107    match type_name {
108        TypeName::Simple(name) => name.to_string(),
109        TypeName::Generic { name, .. } => name.to_string(),
110    }
111}
112
113/// Collect every `impl Trait for Type` block whose target type matches
114/// `type_name`, searching the current program and (optionally) every
115/// importable module reachable from `current_file`.
116///
117/// Returns a deduplicated list ordered by `(trait_name, impl_name,
118/// source_module)` so the resulting markdown render is stable across runs.
119pub fn collect_impls_for_type(
120    program: &Program,
121    type_name: &str,
122    module_cache: Option<&ModuleCache>,
123    current_file: Option<&Path>,
124    workspace_root: Option<&Path>,
125) -> Vec<ImplSummary> {
126    let mut results: Vec<ImplSummary> = Vec::new();
127
128    for item in &program.items {
129        if let Item::Impl(impl_block, _) = item {
130            if type_name_base(&impl_block.target_type) == type_name {
131                results.push(ImplSummary {
132                    trait_name: type_name_base(&impl_block.trait_name),
133                    impl_name: impl_block.impl_name.clone(),
134                    source_module: None,
135                });
136            }
137        }
138    }
139
140    if let (Some(cache), Some(current_file)) = (module_cache, current_file) {
141        let mut import_paths = cache.list_importable_modules_with_context(current_file, workspace_root);
142        import_paths.sort();
143        for import_path in import_paths {
144            let Some(resolved) = cache.resolve_import(&import_path, current_file, workspace_root)
145            else {
146                continue;
147            };
148            let Some(module_info) =
149                cache.load_module_with_context(&resolved, current_file, workspace_root)
150            else {
151                continue;
152            };
153            for item in &module_info.program.items {
154                if let Item::Impl(impl_block, _) = item {
155                    if type_name_base(&impl_block.target_type) == type_name {
156                        results.push(ImplSummary {
157                            trait_name: type_name_base(&impl_block.trait_name),
158                            impl_name: impl_block.impl_name.clone(),
159                            source_module: Some(import_path.clone()),
160                        });
161                    }
162                }
163            }
164        }
165    }
166
167    results.sort_by(|a, b| {
168        a.trait_name
169            .cmp(&b.trait_name)
170            .then_with(|| a.impl_name.cmp(&b.impl_name))
171            .then_with(|| a.source_module.cmp(&b.source_module))
172    });
173    results.dedup_by(|a, b| {
174        a.trait_name == b.trait_name
175            && a.impl_name == b.impl_name
176            && a.source_module == b.source_module
177    });
178
179    results
180}
181
182#[cfg(test)]
183mod tests {
184    use super::*;
185    use shape_ast::parser::parse_program;
186
187    #[test]
188    fn collects_local_impls_for_type() {
189        let program = parse_program(
190            "type User { name: string }\n\
191             impl Display for User { method display() { \"\" } }\n\
192             impl Debug for User as DebugUser { method debug() { \"\" } }\n",
193        )
194        .expect("program");
195
196        let impls = collect_impls_for_type(&program, "User", None, None, None);
197        assert_eq!(impls.len(), 2);
198        assert_eq!(impls[0].trait_name, "Debug");
199        assert_eq!(impls[0].impl_name.as_deref(), Some("DebugUser"));
200        assert_eq!(impls[0].source_module, None);
201        assert_eq!(impls[1].trait_name, "Display");
202        assert_eq!(impls[1].impl_name, None);
203    }
204
205    #[test]
206    fn returns_empty_when_type_has_no_impls() {
207        let program = parse_program("type Lonely { x: int }\n").expect("program");
208        assert!(collect_impls_for_type(&program, "Lonely", None, None, None).is_empty());
209    }
210
211    #[test]
212    fn dedups_identical_local_entries() {
213        let program = parse_program(
214            "type T { x: int }\n\
215             impl Display for T { method display() { \"\" } }\n\
216             impl Display for T { method display() { \"\" } }\n",
217        )
218        .expect("program");
219        let impls = collect_impls_for_type(&program, "T", None, None, None);
220        assert_eq!(impls.len(), 1);
221    }
222}