1use shape_ast::ast::{Expr, Item, Program, Span, Statement};
7use tower_lsp_server::ls_types::{FoldingRange, FoldingRangeKind};
8
9pub fn get_folding_ranges(source: &str, program: &Program) -> Vec<FoldingRange> {
14 let mut ranges = Vec::new();
15
16 collect_comment_folds(source, &mut ranges);
18
19 collect_import_folds(source, program, &mut ranges);
21
22 for item in &program.items {
24 collect_item_folds(source, item, &mut ranges);
25 }
26
27 ranges
28}
29
30fn 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 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 _ => {}
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
175fn 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 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 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
241fn 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 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 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 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}