Skip to main content

shape_lsp/
folding.rs

1//! Folding range support for Shape LSP
2//!
3//! Provides foldable regions for functions, types, traits, impls, enums,
4//! annotations, blocks, and import groups.
5
6use shape_ast::ast::{Expr, Item, Program, Span, Statement};
7use tower_lsp_server::ls_types::{FoldingRange, FoldingRangeKind};
8
9/// Compute folding ranges for a Shape source document.
10///
11/// Walks the AST for multi-line constructs and also scans raw source
12/// for comment blocks and consecutive import groups.
13pub fn get_folding_ranges(source: &str, program: &Program) -> Vec<FoldingRange> {
14    let mut ranges = Vec::new();
15
16    // Collect comment folding ranges from raw source
17    collect_comment_folds(source, &mut ranges);
18
19    // Collect import group folding ranges
20    collect_import_folds(source, program, &mut ranges);
21
22    // Walk AST items for structural folds
23    for item in &program.items {
24        collect_item_folds(source, item, &mut ranges);
25    }
26
27    ranges
28}
29
30/// Convert a byte-offset Span to (start_line, end_line). Returns None if single-line.
31fn span_to_lines(source: &str, span: Span) -> Option<(u32, u32)> {
32    if span.is_empty() || span.is_dummy() {
33        return None;
34    }
35    let start_line = source[..span.start].matches('\n').count() as u32;
36    let end_line = source[..span.end.min(source.len())].matches('\n').count() as u32;
37    if end_line > start_line {
38        Some((start_line, end_line))
39    } else {
40        None
41    }
42}
43
44fn add_region_fold(ranges: &mut Vec<FoldingRange>, start_line: u32, end_line: u32) {
45    ranges.push(FoldingRange {
46        start_line,
47        start_character: None,
48        end_line,
49        end_character: None,
50        kind: Some(FoldingRangeKind::Region),
51        collapsed_text: None,
52    });
53}
54
55fn collect_item_folds(source: &str, item: &Item, ranges: &mut Vec<FoldingRange>) {
56    match item {
57        Item::Function(func, span) => {
58            if let Some((start, end)) = span_to_lines(source, *span) {
59                add_region_fold(ranges, start, end);
60            }
61            // Fold nested blocks in function body
62            for stmt in &func.body {
63                collect_stmt_folds(source, stmt, ranges);
64            }
65        }
66        Item::ForeignFunction(_, span)
67        | Item::StructType(_, span)
68        | Item::Enum(_, span)
69        | Item::Trait(_, span)
70        | Item::Impl(_, span)
71        | Item::Extend(_, span)
72        | Item::AnnotationDef(_, span)
73        | Item::DataSource(_, span)
74        | Item::QueryDecl(_, span)
75        | Item::Stream(_, span)
76        | Item::Test(_, span)
77        | Item::Optimize(_, span) => {
78            if let Some((start, end)) = span_to_lines(source, *span) {
79                add_region_fold(ranges, start, end);
80            }
81        }
82        Item::Statement(stmt, _) => {
83            collect_stmt_folds(source, stmt, ranges);
84        }
85        Item::Expression(expr, _) => {
86            collect_expr_folds(source, expr, ranges);
87        }
88        // Single-line items: imports, exports, variable decls, assignments, comptime
89        _ => {}
90    }
91}
92
93fn collect_stmt_folds(source: &str, stmt: &Statement, ranges: &mut Vec<FoldingRange>) {
94    match stmt {
95        Statement::If(if_stmt, span) => {
96            if let Some((start, end)) = span_to_lines(source, *span) {
97                add_region_fold(ranges, start, end);
98            }
99            for s in &if_stmt.then_body {
100                collect_stmt_folds(source, s, ranges);
101            }
102            if let Some(else_stmts) = &if_stmt.else_body {
103                for s in else_stmts {
104                    collect_stmt_folds(source, s, ranges);
105                }
106            }
107        }
108        Statement::For(_, span) | Statement::While(_, span) => {
109            if let Some((start, end)) = span_to_lines(source, *span) {
110                add_region_fold(ranges, start, end);
111            }
112        }
113        Statement::Expression(expr, _) => {
114            collect_expr_folds(source, expr, ranges);
115        }
116        _ => {}
117    }
118}
119
120fn collect_expr_folds(source: &str, expr: &Expr, ranges: &mut Vec<FoldingRange>) {
121    match expr {
122        Expr::Block(block, span) => {
123            if let Some((start, end)) = span_to_lines(source, *span) {
124                add_region_fold(ranges, start, end);
125            }
126            for item in &block.items {
127                match item {
128                    shape_ast::ast::BlockItem::Statement(s) => {
129                        collect_stmt_folds(source, s, ranges)
130                    }
131                    shape_ast::ast::BlockItem::Expression(e) => {
132                        collect_expr_folds(source, e, ranges)
133                    }
134                    _ => {}
135                }
136            }
137        }
138        Expr::If(if_expr, span) => {
139            if let Some((start, end)) = span_to_lines(source, *span) {
140                add_region_fold(ranges, start, end);
141            }
142            collect_expr_folds(source, &if_expr.then_branch, ranges);
143            if let Some(else_br) = &if_expr.else_branch {
144                collect_expr_folds(source, else_br, ranges);
145            }
146        }
147        Expr::Conditional {
148            span,
149            then_expr,
150            else_expr,
151            ..
152        } => {
153            if let Some((start, end)) = span_to_lines(source, *span) {
154                add_region_fold(ranges, start, end);
155            }
156            collect_expr_folds(source, then_expr, ranges);
157            if let Some(else_br) = else_expr {
158                collect_expr_folds(source, else_br, ranges);
159            }
160        }
161        Expr::For(_, span) | Expr::While(_, span) | Expr::Loop(_, span) | Expr::Match(_, span) => {
162            if let Some((start, end)) = span_to_lines(source, *span) {
163                add_region_fold(ranges, start, end);
164            }
165        }
166        Expr::FunctionExpr { body, .. } => {
167            for stmt in body {
168                collect_stmt_folds(source, stmt, ranges);
169            }
170        }
171        _ => {}
172    }
173}
174
175/// Scan source for consecutive line comments (// ...) and block comments (/* ... */).
176fn collect_comment_folds(source: &str, ranges: &mut Vec<FoldingRange>) {
177    let lines: Vec<&str> = source.lines().collect();
178    let mut i = 0;
179    while i < lines.len() {
180        let trimmed = lines[i].trim_start();
181        // Consecutive line comments
182        if trimmed.starts_with("//") {
183            let start = i;
184            while i < lines.len() && lines[i].trim_start().starts_with("//") {
185                i += 1;
186            }
187            let end = i - 1;
188            if end > start {
189                ranges.push(FoldingRange {
190                    start_line: start as u32,
191                    start_character: None,
192                    end_line: end as u32,
193                    end_character: None,
194                    kind: Some(FoldingRangeKind::Comment),
195                    collapsed_text: None,
196                });
197            }
198            continue;
199        }
200        // Block comments: find /* and scan to */
201        if trimmed.starts_with("/*") {
202            let start = i;
203            let mut depth = 0u32;
204            let mut found_end = false;
205            while i < lines.len() {
206                let line = lines[i];
207                for (idx, _) in line.char_indices() {
208                    if line[idx..].starts_with("/*") {
209                        depth += 1;
210                    } else if line[idx..].starts_with("*/") {
211                        depth = depth.saturating_sub(1);
212                        if depth == 0 {
213                            found_end = true;
214                            break;
215                        }
216                    }
217                }
218                if found_end {
219                    break;
220                }
221                i += 1;
222            }
223            let end = i;
224            if end > start {
225                ranges.push(FoldingRange {
226                    start_line: start as u32,
227                    start_character: None,
228                    end_line: end as u32,
229                    end_character: None,
230                    kind: Some(FoldingRangeKind::Comment),
231                    collapsed_text: None,
232                });
233            }
234            i += 1;
235            continue;
236        }
237        i += 1;
238    }
239}
240
241/// Group consecutive import statements into a single Imports fold.
242fn collect_import_folds(source: &str, program: &Program, ranges: &mut Vec<FoldingRange>) {
243    let mut import_lines: Vec<u32> = Vec::new();
244    for item in &program.items {
245        if let Item::Import(_, span) = item {
246            if !span.is_dummy() {
247                let line = source[..span.start].matches('\n').count() as u32;
248                import_lines.push(line);
249            }
250        }
251    }
252    if import_lines.len() < 2 {
253        return;
254    }
255    import_lines.sort();
256
257    // Group consecutive lines (allowing gaps of 1 blank line)
258    let mut group_start = import_lines[0];
259    let mut group_end = import_lines[0];
260    for &line in &import_lines[1..] {
261        if line <= group_end + 2 {
262            group_end = line;
263        } else {
264            if group_end > group_start {
265                ranges.push(FoldingRange {
266                    start_line: group_start,
267                    start_character: None,
268                    end_line: group_end,
269                    end_character: None,
270                    kind: Some(FoldingRangeKind::Imports),
271                    collapsed_text: None,
272                });
273            }
274            group_start = line;
275            group_end = line;
276        }
277    }
278    if group_end > group_start {
279        ranges.push(FoldingRange {
280            start_line: group_start,
281            start_character: None,
282            end_line: group_end,
283            end_character: None,
284            kind: Some(FoldingRangeKind::Imports),
285            collapsed_text: None,
286        });
287    }
288}
289
290#[cfg(test)]
291mod tests {
292    use super::*;
293    use shape_ast::parser::parse_program;
294
295    fn fold_kinds(source: &str) -> Vec<(u32, u32, Option<FoldingRangeKind>)> {
296        let program = parse_program(source).expect("parse should succeed");
297        let ranges = get_folding_ranges(source, &program);
298        ranges
299            .into_iter()
300            .map(|r| (r.start_line, r.end_line, r.kind))
301            .collect()
302    }
303
304    #[test]
305    fn test_function_fold() {
306        let source = "fn foo(a) {\n  return a\n}";
307        let folds = fold_kinds(source);
308        assert!(
309            folds
310                .iter()
311                .any(|(s, e, k)| *s == 0 && *e == 2 && *k == Some(FoldingRangeKind::Region)),
312            "expected function fold 0..2, got: {:?}",
313            folds
314        );
315    }
316
317    #[test]
318    fn test_enum_fold() {
319        let source = "enum Color {\n  Red,\n  Green,\n  Blue\n}";
320        let folds = fold_kinds(source);
321        assert!(
322            folds
323                .iter()
324                .any(|(s, e, k)| *s == 0 && *e == 4 && *k == Some(FoldingRangeKind::Region)),
325            "expected enum fold 0..4, got: {:?}",
326            folds
327        );
328    }
329
330    #[test]
331    fn test_trait_fold() {
332        let source = "trait Printable {\n  method to_string() -> string {\n    return \"\"\n  }\n}";
333        let folds = fold_kinds(source);
334        assert!(
335            folds
336                .iter()
337                .any(|(s, _e, k)| *s == 0 && *k == Some(FoldingRangeKind::Region)),
338            "expected trait fold starting at line 0, got: {:?}",
339            folds
340        );
341    }
342
343    #[test]
344    fn test_comment_fold() {
345        let source = "// line 1\n// line 2\n// line 3\nlet x = 1";
346        let program = parse_program(source).expect("parse");
347        let ranges = get_folding_ranges(source, &program);
348        assert!(
349            ranges.iter().any(|r| r.start_line == 0
350                && r.end_line == 2
351                && r.kind == Some(FoldingRangeKind::Comment)),
352            "expected comment fold 0..2, got: {:?}",
353            ranges
354        );
355    }
356
357    #[test]
358    fn test_import_fold() {
359        let source = "from a use { a }\nfrom b use { b }\nfrom c use { c }\nlet x = 1";
360        let program = parse_program(source).expect("parse");
361        let ranges = get_folding_ranges(source, &program);
362        assert!(
363            ranges
364                .iter()
365                .any(|r| r.kind == Some(FoldingRangeKind::Imports)),
366            "expected import fold, got: {:?}",
367            ranges
368        );
369    }
370
371    #[test]
372    fn test_single_line_no_fold() {
373        let source = "let x = 1";
374        let program = parse_program(source).expect("parse");
375        let ranges = get_folding_ranges(source, &program);
376        // Should have no region folds for single-line content
377        assert!(
378            !ranges
379                .iter()
380                .any(|r| r.kind == Some(FoldingRangeKind::Region)),
381            "single line should not produce region folds, got: {:?}",
382            ranges
383        );
384    }
385
386    #[test]
387    fn test_nested_folds() {
388        let source = "fn foo() {\n  if true {\n    let x = 1\n  }\n}";
389        let folds = fold_kinds(source);
390        // Should have at least 2 region folds (function + if block)
391        let region_folds: Vec<_> = folds
392            .iter()
393            .filter(|(_, _, k)| *k == Some(FoldingRangeKind::Region))
394            .collect();
395        assert!(
396            region_folds.len() >= 2,
397            "expected at least 2 nested folds, got: {:?}",
398            region_folds
399        );
400    }
401}