1use shape_ast::ast::{Item, Program};
6use shape_ast::parser::parse_program;
7use tower_lsp_server::ls_types::{CodeLens, Command, Position, Range, Uri};
8
9pub fn get_code_lenses(text: &str, uri: &Uri) -> Vec<CodeLens> {
11 let mut lenses = Vec::new();
12
13 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 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
37pub fn resolve_code_lens(lens: CodeLens) -> CodeLens {
39 lens
41}
42
43fn 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 if let Some((line, keyword_end_col)) = find_function_line(text, &func.name) {
56 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 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 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 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 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 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
211fn 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
227fn 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 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#[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
257fn 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
268fn 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
276fn 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 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
308fn 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 let count = binding.references.len();
324 total = Some(total.map_or(count, |t| t + count));
325 }
326 }
327 total
328}
329
330fn 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 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 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 assert_eq!(count_references(text, "foo"), 2);
365
366 assert_eq!(count_references(text, "bar"), 0);
368
369 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 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 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 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 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 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 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}