Skip to main content

shape_lsp/
code_lens.rs

1//! Code lens provider for Shape
2//!
3//! Provides actionable code lenses for functions, patterns, and tests.
4
5use shape_ast::ast::{Item, Program};
6use shape_ast::parser::parse_program;
7use tower_lsp_server::ls_types::{CodeLens, Command, Position, Range, Uri};
8
9/// Get code lenses for a document
10pub fn get_code_lenses(text: &str, uri: &Uri) -> Vec<CodeLens> {
11    let mut lenses = Vec::new();
12
13    // Parse the document, falling back to resilient parser
14    let program = match parse_program(text) {
15        Ok(p) => p,
16        Err(_) => {
17            let partial = shape_ast::parse_program_resilient(text);
18            if partial.items.is_empty() {
19                return lenses;
20            }
21            partial.into_program()
22        }
23    };
24
25    // W2.6 1.55 — build ScopeTree once per document so per-function ref
26    // counts use scope-aware resolution (not text-search) — same ref count
27    // as `references_provider` produces for the same symbol in-file.
28    let tree = crate::scope::ScopeTree::build(&program, text);
29
30    for item in &program.items {
31        collect_lenses_for_item(item, &program, &tree, text, uri, &mut lenses);
32    }
33
34    lenses
35}
36
37/// Resolve a code lens (add the command)
38pub fn resolve_code_lens(lens: CodeLens) -> CodeLens {
39    // Code lenses are already resolved in get_code_lenses
40    lens
41}
42
43/// Collect code lenses for an item
44fn collect_lenses_for_item(
45    item: &Item,
46    _program: &Program,
47    tree: &crate::scope::ScopeTree,
48    text: &str,
49    uri: &Uri,
50    lenses: &mut Vec<CodeLens>,
51) {
52    match item {
53        Item::Function(func, _) => {
54            // Find the line where the function is defined
55            if let Some((line, keyword_end_col)) = find_function_line(text, &func.name) {
56                // W2.6 1.55 — scope-aware ref count from module scope
57                // (top-level fn bindings live there). Falls back to
58                // text-search if ScopeTree didn't record this binding
59                // (e.g. resilient-parse partial AST).
60                let ref_count = count_references_scope_aware(tree, &func.name)
61                    .unwrap_or_else(|| count_references(text, &func.name));
62                lenses.push(CodeLens {
63                    range: Range {
64                        start: Position { line, character: 0 },
65                        end: Position { line, character: 0 },
66                    },
67                    command: Some(Command {
68                        title: format!(
69                            "{} reference{}",
70                            ref_count,
71                            if ref_count == 1 { "" } else { "s" }
72                        ),
73                        command: "shape.findReferences".to_string(),
74                        arguments: Some(vec![
75                            serde_json::json!(uri.to_string()),
76                            serde_json::json!(line),
77                            serde_json::json!(keyword_end_col),
78                        ]),
79                    }),
80                    data: None,
81                });
82
83                // Add code lenses for annotations on self function
84                for annotation in &func.annotations {
85                    lenses.push(CodeLens {
86                        range: Range {
87                            start: Position { line, character: 0 },
88                            end: Position { line, character: 0 },
89                        },
90                        command: Some(Command {
91                            title: format!("@{}", annotation.name),
92                            command: "shape.showAnnotation".to_string(),
93                            arguments: Some(vec![
94                                serde_json::json!(uri.to_string()),
95                                serde_json::json!(annotation.name),
96                                serde_json::json!(func.name),
97                            ]),
98                        }),
99                        data: None,
100                    });
101                }
102            }
103        }
104        Item::Trait(trait_def, _) => {
105            // Add "N implementations" lens on the trait definition
106            if let Some(line) = find_trait_line(text, &trait_def.name) {
107                let impl_count = count_trait_implementations(text, &trait_def.name);
108                lenses.push(CodeLens {
109                    range: Range {
110                        start: Position { line, character: 0 },
111                        end: Position { line, character: 0 },
112                    },
113                    command: Some(Command {
114                        title: format!(
115                            "{} implementation{}",
116                            impl_count,
117                            if impl_count == 1 { "" } else { "s" }
118                        ),
119                        command: "shape.findImplementations".to_string(),
120                        arguments: Some(vec![
121                            serde_json::json!(uri.to_string()),
122                            serde_json::json!(trait_def.name),
123                        ]),
124                    }),
125                    data: None,
126                });
127            }
128
129            // Add per-method lenses showing if the method has a default implementation
130            for member in &trait_def.members {
131                let (method_name, is_default) = match member {
132                    shape_ast::ast::TraitMember::Required(
133                        shape_ast::ast::TraitMemberSignature::Method { name, .. },
134                    ) => (name.as_str(), false),
135                    shape_ast::ast::TraitMember::Default(method_def) => {
136                        (method_def.name.as_str(), true)
137                    }
138                    _ => continue,
139                };
140
141                if let Some(method_line) = find_method_in_trait(text, &trait_def.name, method_name)
142                {
143                    if is_default {
144                        lenses.push(CodeLens {
145                            range: Range {
146                                start: Position {
147                                    line: method_line,
148                                    character: 0,
149                                },
150                                end: Position {
151                                    line: method_line,
152                                    character: 0,
153                                },
154                            },
155                            command: Some(Command {
156                                title: "(default)".to_string(),
157                                command: "shape.showTraitMethod".to_string(),
158                                arguments: Some(vec![
159                                    serde_json::json!(uri.to_string()),
160                                    serde_json::json!(trait_def.name),
161                                    serde_json::json!(method_name),
162                                ]),
163                            }),
164                            data: None,
165                        });
166                    }
167                }
168            }
169        }
170        Item::Test(test, _) => {
171            if let Some(line) = find_test_line(text, &test.name) {
172                // Run all tests lens
173                lenses.push(CodeLens {
174                    range: Range {
175                        start: Position { line, character: 0 },
176                        end: Position { line, character: 0 },
177                    },
178                    command: Some(Command {
179                        title: "▶ Run All Tests".to_string(),
180                        command: "shape.runTests".to_string(),
181                        arguments: Some(vec![
182                            serde_json::json!(uri.to_string()),
183                            serde_json::json!(test.name),
184                        ]),
185                    }),
186                    data: None,
187                });
188
189                // Debug tests lens
190                lenses.push(CodeLens {
191                    range: Range {
192                        start: Position { line, character: 0 },
193                        end: Position { line, character: 0 },
194                    },
195                    command: Some(Command {
196                        title: "🐛 Debug Tests".to_string(),
197                        command: "shape.debugTests".to_string(),
198                        arguments: Some(vec![
199                            serde_json::json!(uri.to_string()),
200                            serde_json::json!(test.name),
201                        ]),
202                    }),
203                    data: None,
204                });
205            }
206        }
207        _ => {}
208    }
209}
210
211/// Find the line number where a function is defined
212fn find_function_line(text: &str, name: &str) -> Option<(u32, u32)> {
213    let fn_pattern = format!("fn {}", name);
214    let function_pattern = format!("function {}", name);
215
216    for (line_num, line) in text.lines().enumerate() {
217        if let Some(col) = line.find(&fn_pattern) {
218            return Some((line_num as u32, (col + "fn ".len()) as u32));
219        }
220        if let Some(col) = line.find(&function_pattern) {
221            return Some((line_num as u32, (col + "function ".len()) as u32));
222        }
223    }
224    None
225}
226
227/// Find the line number where a test is defined
228fn find_test_line(text: &str, name: &str) -> Option<u32> {
229    let pattern = format!("test \"{}\"", name);
230    for (line_num, line) in text.lines().enumerate() {
231        if line.contains(&pattern) {
232            return Some(line_num as u32);
233        }
234    }
235    // Also try without quotes
236    let pattern = format!("test {}", name);
237    for (line_num, line) in text.lines().enumerate() {
238        if line.contains(&pattern) {
239            return Some(line_num as u32);
240        }
241    }
242    None
243}
244
245/// Find the line number where a pattern is defined
246#[allow(dead_code)]
247fn find_pattern_line(text: &str, name: &str) -> Option<u32> {
248    let pattern = format!("pattern {}", name);
249    for (line_num, line) in text.lines().enumerate() {
250        if line.contains(&pattern) {
251            return Some(line_num as u32);
252        }
253    }
254    None
255}
256
257/// Find the line number where a trait is defined
258fn find_trait_line(text: &str, name: &str) -> Option<u32> {
259    let pattern = format!("trait {}", name);
260    for (line_num, line) in text.lines().enumerate() {
261        if line.trim().starts_with(&pattern) {
262            return Some(line_num as u32);
263        }
264    }
265    None
266}
267
268/// Count the number of `impl TraitName for ...` blocks in the text
269fn count_trait_implementations(text: &str, trait_name: &str) -> usize {
270    let pattern = format!("impl {} for", trait_name);
271    text.lines()
272        .filter(|line| line.trim().starts_with(&pattern) || line.trim().contains(&pattern))
273        .count()
274}
275
276/// Find the line of a method within a trait definition
277fn find_method_in_trait(text: &str, trait_name: &str, method_name: &str) -> Option<u32> {
278    let trait_pattern = format!("trait {}", trait_name);
279    let mut in_trait = false;
280    let mut brace_count: i32 = 0;
281
282    for (line_num, line) in text.lines().enumerate() {
283        if line.trim().starts_with(&trait_pattern) {
284            in_trait = true;
285        }
286
287        if in_trait {
288            brace_count += line.matches('{').count() as i32;
289            brace_count -= line.matches('}').count() as i32;
290
291            // Check if self line contains the method name
292            let trimmed = line.trim();
293            if (trimmed.contains(&format!("{}(", method_name))
294                || trimmed.starts_with(&format!("method {}(", method_name)))
295                && !trimmed.starts_with("trait ")
296            {
297                return Some(line_num as u32);
298            }
299
300            if brace_count == 0 && line.contains('}') {
301                in_trait = false;
302            }
303        }
304    }
305    None
306}
307
308/// W2.6 1.55 — scope-aware reference count for the module-scope binding
309/// `name` in `tree`. Returns `None` when no module-scope binding with this
310/// name exists (caller falls back to text-search for robustness on
311/// resilient-parse partials).
312fn count_references_scope_aware(
313    tree: &crate::scope::ScopeTree,
314    name: &str,
315) -> Option<usize> {
316    let root = tree.scopes.first()?;
317    let mut total: Option<usize> = None;
318    for binding in &root.bindings {
319        if binding.name == name {
320            // ScopeTree::Binding::references excludes the def site itself,
321            // matching the legacy text-search semantics (which counted
322            // def+uses then subtracted 1).
323            let count = binding.references.len();
324            total = Some(total.map_or(count, |t| t + count));
325        }
326    }
327    total
328}
329
330/// Count references to a symbol in the text
331fn count_references(text: &str, name: &str) -> usize {
332    let mut count = 0;
333    let name_len = name.len();
334
335    for (i, _) in text.match_indices(name) {
336        // Check word boundaries
337        let before_ok = i == 0 || !text[..i].chars().last().unwrap().is_alphanumeric();
338        let after_ok = i + name_len >= text.len()
339            || !text[i + name_len..]
340                .chars()
341                .next()
342                .unwrap()
343                .is_alphanumeric();
344
345        if before_ok && after_ok {
346            count += 1;
347        }
348    }
349
350    // Subtract 1 for the definition itself
351    if count > 0 { count - 1 } else { 0 }
352}
353
354#[cfg(test)]
355mod tests {
356    use super::*;
357
358    #[test]
359    fn test_count_references() {
360        let text = "let foo = 1;\nlet bar = foo + foo;";
361
362        // foo appears 3 times (1 definition + 2 uses)
363        // count_references subtracts 1 for definition
364        assert_eq!(count_references(text, "foo"), 2);
365
366        // bar appears once (just definition)
367        assert_eq!(count_references(text, "bar"), 0);
368
369        // nonexistent
370        assert_eq!(count_references(text, "baz"), 0);
371    }
372
373    #[test]
374    fn test_find_function_line() {
375        let text = "// comment\nfunction myFunc() {\n    return 1;\n}";
376        assert_eq!(find_function_line(text, "myFunc"), Some((1, 9)));
377        let text = "// comment\nfn myFunc() {\n    return 1;\n}";
378        assert_eq!(find_function_line(text, "myFunc"), Some((1, 3)));
379        assert_eq!(find_function_line(text, "nonexistent"), None);
380    }
381
382    #[test]
383    fn test_find_trait_line() {
384        let text = "// comment\ntrait Queryable {\n    filter(pred): any\n}\n";
385        assert_eq!(find_trait_line(text, "Queryable"), Some(1));
386        assert_eq!(find_trait_line(text, "NonExistent"), None);
387    }
388
389    #[test]
390    fn test_count_trait_implementations() {
391        let text = "trait Queryable {\n    filter(pred): any\n}\nimpl Queryable for Table {\n    method filter(pred) { self }\n}\nimpl Queryable for DataFrame {\n    method filter(pred) { self }\n}\n";
392        assert_eq!(count_trait_implementations(text, "Queryable"), 2);
393        assert_eq!(count_trait_implementations(text, "NonExistent"), 0);
394    }
395
396    #[test]
397    fn test_trait_code_lens() {
398        // Fixture migrated to Form B per cd7d97a4 (2026-05-18) grammar surgery —
399        // Form A `name(params): RetType` was deleted from trait_member_signature.
400        let text = "trait Queryable {\n    method filter(self, pred) -> any;\n}\nimpl Queryable for Table {\n    method filter(pred) { self }\n}\n";
401        let uri = Uri::from_file_path("/tmp/test.shape").unwrap();
402        let lenses = get_code_lenses(text, &uri);
403        // Should have at least one code lens for the trait
404        assert!(
405            lenses.iter().any(|l| l
406                .command
407                .as_ref()
408                .map_or(false, |c| c.title.contains("implementation"))),
409            "Should have implementation count lens for trait. Got: {:?}",
410            lenses
411                .iter()
412                .map(|l| l.command.as_ref().map(|c| c.title.clone()))
413                .collect::<Vec<_>>()
414        );
415    }
416
417    #[test]
418    fn test_count_references_scope_aware_excludes_shadowing() {
419        // W2.6 1.55 — ScopeTree-based count must NOT include shadowing inner
420        // bindings as references to the outer one. This is the bug the
421        // text-search count had: a shadowed inner `let foo` was counted.
422        let text = "fn foo() { return 1 }\nfn other() {\n  let foo = 2\n  return foo + foo\n}\nlet x = foo()";
423        let program = parse_program(text).unwrap();
424        let tree = crate::scope::ScopeTree::build(&program, text);
425
426        let scope_count = count_references_scope_aware(&tree, "foo");
427        // Top-level `foo` (the fn) is referenced exactly once: `foo()` at the
428        // end. The `let foo = 2` and `foo + foo` inside `other` shadow.
429        assert_eq!(
430            scope_count,
431            Some(1),
432            "expected scope-aware count to exclude shadowing inner foo, got {:?}",
433            scope_count
434        );
435    }
436
437    #[test]
438    fn test_count_references_scope_aware_lens_integration() {
439        // End-to-end: get_code_lenses must produce a "N references" lens
440        // whose count matches the scope-aware count.
441        let text = "fn helper() { return 1 }\nlet a = helper()\nlet b = helper() + helper()";
442        let uri = Uri::from_file_path("/tmp/test.shape").unwrap();
443        let lenses = get_code_lenses(text, &uri);
444        let helper_lens = lenses
445            .iter()
446            .find(|l| {
447                l.command
448                    .as_ref()
449                    .is_some_and(|c| c.title.contains("reference"))
450            })
451            .expect("should have reference-count lens for helper");
452        let title = &helper_lens.command.as_ref().unwrap().title;
453        // 3 references: a = helper(), b = helper() + helper()
454        assert!(
455            title.starts_with("3 references"),
456            "expected '3 references' from scope-aware count, got '{}'",
457            title
458        );
459    }
460
461    #[test]
462    fn test_find_pattern_line() {
463        let text = "// comment\npattern hammer {\n    close > open\n}";
464        assert_eq!(find_pattern_line(text, "hammer"), Some(1));
465        assert_eq!(find_pattern_line(text, "doji"), None);
466    }
467}