1use ryo_source::pure::{PureBlock, PureClosureParam, PureExpr, PurePattern, PureStmt, PureType};
17use ryo_symbol::SymbolId;
18
19use crate::Mutation;
20
21#[derive(Debug, Clone, PartialEq)]
23pub enum LoopPattern {
24 MapCollect {
26 iter_expr: PureExpr,
27 var_name: String,
28 target_var: String,
29 transform: PureExpr,
30 },
31 FilterCollect {
33 iter_expr: PureExpr,
34 var_name: String,
35 target_var: String,
36 condition: PureExpr,
37 },
38 FilterMapCollect {
40 iter_expr: PureExpr,
41 var_name: String,
42 target_var: String,
43 condition: PureExpr,
44 transform: PureExpr,
45 },
46 ForEach {
48 iter_expr: PureExpr,
49 var_name: String,
50 body: PureBlock,
51 },
52}
53
54#[derive(Debug, Clone)]
56pub struct BlockLoopPattern {
57 pub vec_decl_idx: usize,
59 pub for_loop_idx: usize,
61 pub target_var: String,
63 pub loop_pattern: LoopPattern,
65}
66
67#[derive(Debug, Clone, Default)]
69pub struct LoopToIteratorMutation {
70 pub target_var: Option<String>,
72 pub aggressive: bool,
74 pub target_fn: Option<SymbolId>,
76}
77
78impl LoopToIteratorMutation {
79 pub fn new() -> Self {
80 Self::default()
81 }
82
83 pub fn with_target(mut self, var: impl Into<String>) -> Self {
85 self.target_var = Some(var.into());
86 self
87 }
88
89 pub fn aggressive(mut self) -> Self {
91 self.aggressive = true;
92 self
93 }
94
95 pub fn in_function(mut self, id: SymbolId) -> Self {
97 self.target_fn = Some(id);
98 self
99 }
100
101 fn detect_pattern(
103 var_pattern: &PurePattern,
104 iter_expr: &PureExpr,
105 body: &PureBlock,
106 ) -> Option<LoopPattern> {
107 let var_name = match var_pattern {
108 PurePattern::Ident { name, .. } => name.clone(),
109 _ => return None, };
111
112 if body.stmts.len() == 1 {
114 if let Some(pattern) =
115 Self::detect_single_stmt_pattern(&var_name, iter_expr, &body.stmts[0])
116 {
117 return Some(pattern);
118 }
119 }
120
121 if body.stmts.len() == 1 {
123 if let PureStmt::Semi(PureExpr::If {
124 cond,
125 then_branch,
126 else_branch: None,
127 })
128 | PureStmt::Expr(PureExpr::If {
129 cond,
130 then_branch,
131 else_branch: None,
132 }) = &body.stmts[0]
133 {
134 if then_branch.stmts.len() == 1 {
135 if let Some((target_var, pushed_expr)) =
136 Self::extract_push(&then_branch.stmts[0])
137 {
138 if Self::is_simple_var(&pushed_expr, &var_name) {
140 return Some(LoopPattern::FilterCollect {
141 iter_expr: iter_expr.clone(),
142 var_name,
143 target_var,
144 condition: *cond.clone(),
145 });
146 } else {
147 return Some(LoopPattern::FilterMapCollect {
148 iter_expr: iter_expr.clone(),
149 var_name,
150 target_var,
151 condition: *cond.clone(),
152 transform: pushed_expr,
153 });
154 }
155 }
156 }
157 }
158 }
159
160 Some(LoopPattern::ForEach {
162 iter_expr: iter_expr.clone(),
163 var_name,
164 body: body.clone(),
165 })
166 }
167
168 fn detect_single_stmt_pattern(
170 var_name: &str,
171 iter_expr: &PureExpr,
172 stmt: &PureStmt,
173 ) -> Option<LoopPattern> {
174 if let Some((target_var, pushed_expr)) = Self::extract_push(stmt) {
176 if Self::is_simple_var(&pushed_expr, var_name) {
177 return Some(LoopPattern::MapCollect {
179 iter_expr: iter_expr.clone(),
180 var_name: var_name.to_string(),
181 target_var,
182 transform: pushed_expr,
183 });
184 } else {
185 return Some(LoopPattern::MapCollect {
187 iter_expr: iter_expr.clone(),
188 var_name: var_name.to_string(),
189 target_var,
190 transform: pushed_expr,
191 });
192 }
193 }
194 None
195 }
196
197 fn extract_push(stmt: &PureStmt) -> Option<(String, PureExpr)> {
199 let expr = match stmt {
200 PureStmt::Semi(e) | PureStmt::Expr(e) => e,
201 _ => return None,
202 };
203
204 if let PureExpr::MethodCall {
206 receiver,
207 method,
208 args,
209 ..
210 } = expr
211 {
212 if method == "push" && args.len() == 1 {
213 if let PureExpr::Path(target_var) = receiver.as_ref() {
214 return Some((target_var.clone(), args[0].clone()));
215 }
216 }
217 }
218 None
219 }
220
221 fn is_simple_var(expr: &PureExpr, var_name: &str) -> bool {
223 matches!(expr, PureExpr::Path(name) if name == var_name)
224 }
225
226 fn pattern_to_iter_expr(pattern: &LoopPattern) -> PureExpr {
228 match pattern {
229 LoopPattern::MapCollect {
230 iter_expr,
231 var_name,
232 transform,
233 ..
234 } => {
235 let is_identity = Self::is_simple_var(transform, var_name);
236
237 if is_identity {
238 PureExpr::MethodCall {
240 receiver: Box::new(PureExpr::MethodCall {
241 receiver: Box::new(iter_expr.clone()),
242 method: "into_iter".to_string(),
243 turbofish: None,
244 args: vec![],
245 }),
246 method: "collect".to_string(),
247 turbofish: None,
248 args: vec![],
249 }
250 } else {
251 PureExpr::MethodCall {
253 receiver: Box::new(PureExpr::MethodCall {
254 receiver: Box::new(PureExpr::MethodCall {
255 receiver: Box::new(iter_expr.clone()),
256 method: "into_iter".to_string(),
257 turbofish: None,
258 args: vec![],
259 }),
260 method: "map".to_string(),
261 turbofish: None,
262 args: vec![PureExpr::Closure {
263 is_async: false,
264 is_move: false,
265 params: vec![PureClosureParam::untyped(PurePattern::Ident {
266 name: var_name.clone(),
267 is_mut: false,
268 by_ref: false,
269 })],
270 ret: None,
271 body: Box::new(transform.clone()),
272 }],
273 }),
274 method: "collect".to_string(),
275 turbofish: None,
276 args: vec![],
277 }
278 }
279 }
280 LoopPattern::FilterCollect {
281 iter_expr,
282 var_name,
283 condition,
284 ..
285 } => {
286 PureExpr::MethodCall {
290 receiver: Box::new(PureExpr::MethodCall {
291 receiver: Box::new(PureExpr::MethodCall {
292 receiver: Box::new(iter_expr.clone()),
293 method: "into_iter".to_string(),
294 turbofish: None,
295 args: vec![],
296 }),
297 method: "filter".to_string(),
298 turbofish: None,
299 args: vec![PureExpr::Closure {
300 is_async: false,
301 is_move: false,
302 params: vec![PureClosureParam::untyped(PurePattern::Ref {
303 is_mut: false,
304 pattern: Box::new(PurePattern::Ident {
305 name: var_name.clone(),
306 is_mut: false,
307 by_ref: false,
308 }),
309 })],
310 ret: None,
311 body: Box::new(condition.clone()),
312 }],
313 }),
314 method: "collect".to_string(),
315 turbofish: None,
316 args: vec![],
317 }
318 }
319 LoopPattern::FilterMapCollect {
320 iter_expr,
321 var_name,
322 condition,
323 transform,
324 ..
325 } => {
326 PureExpr::MethodCall {
330 receiver: Box::new(PureExpr::MethodCall {
331 receiver: Box::new(PureExpr::MethodCall {
332 receiver: Box::new(PureExpr::MethodCall {
333 receiver: Box::new(iter_expr.clone()),
334 method: "into_iter".to_string(),
335 turbofish: None,
336 args: vec![],
337 }),
338 method: "filter".to_string(),
339 turbofish: None,
340 args: vec![PureExpr::Closure {
341 is_async: false,
342 is_move: false,
343 params: vec![PureClosureParam::untyped(PurePattern::Ref {
344 is_mut: false,
345 pattern: Box::new(PurePattern::Ident {
346 name: var_name.clone(),
347 is_mut: false,
348 by_ref: false,
349 }),
350 })],
351 ret: None,
352 body: Box::new(condition.clone()),
353 }],
354 }),
355 method: "map".to_string(),
356 turbofish: None,
357 args: vec![PureExpr::Closure {
358 is_async: false,
359 is_move: false,
360 params: vec![PureClosureParam::untyped(PurePattern::Ident {
361 name: var_name.clone(),
362 is_mut: false,
363 by_ref: false,
364 })],
365 ret: None,
366 body: Box::new(transform.clone()),
367 }],
368 }),
369 method: "collect".to_string(),
370 turbofish: None,
371 args: vec![],
372 }
373 }
374 LoopPattern::ForEach {
375 iter_expr,
376 var_name,
377 body,
378 } => {
379 PureExpr::MethodCall {
381 receiver: Box::new(PureExpr::MethodCall {
382 receiver: Box::new(iter_expr.clone()),
383 method: "into_iter".to_string(),
384 turbofish: None,
385 args: vec![],
386 }),
387 method: "for_each".to_string(),
388 turbofish: None,
389 args: vec![PureExpr::Closure {
390 is_async: false,
391 is_move: false,
392 params: vec![PureClosureParam::untyped(PurePattern::Ident {
393 name: var_name.clone(),
394 is_mut: false,
395 by_ref: false,
396 })],
397 ret: None,
398 body: Box::new(PureExpr::Block {
399 label: None,
400 block: body.clone(),
401 }),
402 }],
403 }
404 }
405 }
406 }
407
408 fn detect_block_patterns(stmts: &[PureStmt]) -> Vec<BlockLoopPattern> {
410 let mut patterns = Vec::new();
411
412 for (i, stmt) in stmts.iter().enumerate() {
413 if let Some((var_name, is_vec_init)) = Self::extract_vec_init(stmt) {
415 if !is_vec_init {
416 continue;
417 }
418
419 for (j, jstmt) in stmts.iter().enumerate().skip(i + 1) {
421 if let Some(loop_pattern) = Self::extract_for_loop_push(jstmt, &var_name) {
422 patterns.push(BlockLoopPattern {
423 vec_decl_idx: i,
424 for_loop_idx: j,
425 target_var: var_name.clone(),
426 loop_pattern,
427 });
428 break; }
430
431 if Self::stmt_uses_var(jstmt, &var_name) {
433 break;
434 }
435 }
436 }
437 }
438
439 patterns
440 }
441
442 fn extract_vec_init(stmt: &PureStmt) -> Option<(String, bool)> {
444 if let PureStmt::Local {
445 pattern: PurePattern::Ident {
446 name, is_mut: true, ..
447 },
448 init: Some(init_expr),
449 ..
450 } = stmt
451 {
452 let is_vec_init = Self::is_vec_new_call(init_expr) || Self::is_vec_macro(init_expr);
453 return Some((name.clone(), is_vec_init));
454 }
455 None
456 }
457
458 fn is_vec_new_call(expr: &PureExpr) -> bool {
460 if let PureExpr::Call { func, args } = expr {
461 if args.is_empty() {
462 if let PureExpr::Path(path) = func.as_ref() {
464 return path == "Vec::new" || path.ends_with("::Vec::new");
465 }
466 }
467 }
468 false
469 }
470
471 fn is_vec_macro(expr: &PureExpr) -> bool {
473 if let PureExpr::Macro { name, .. } = expr {
474 return name == "vec" || name.ends_with("::vec");
475 }
476 false
477 }
478
479 fn extract_for_loop_push(stmt: &PureStmt, target_var: &str) -> Option<LoopPattern> {
481 let for_expr = match stmt {
482 PureStmt::Semi(e) | PureStmt::Expr(e) => e,
483 _ => return None,
484 };
485
486 if let PureExpr::For {
487 pat,
488 expr: iter_expr,
489 body,
490 ..
491 } = for_expr
492 {
493 let pattern = Self::detect_pattern(pat, iter_expr, body)?;
494
495 let pattern_target = match &pattern {
497 LoopPattern::MapCollect { target_var, .. } => target_var,
498 LoopPattern::FilterCollect { target_var, .. } => target_var,
499 LoopPattern::FilterMapCollect { target_var, .. } => target_var,
500 LoopPattern::ForEach { .. } => return None, };
502
503 if pattern_target == target_var {
504 return Some(pattern);
505 }
506 }
507 None
508 }
509
510 fn stmt_uses_var(stmt: &PureStmt, var_name: &str) -> bool {
512 match stmt {
513 PureStmt::Semi(expr) | PureStmt::Expr(expr) => Self::expr_uses_var(expr, var_name),
514 PureStmt::Local { init: Some(e), .. } => Self::expr_uses_var(e, var_name),
515 _ => false,
516 }
517 }
518
519 fn expr_uses_var(expr: &PureExpr, var_name: &str) -> bool {
521 match expr {
522 PureExpr::Path(name) => name == var_name,
523 PureExpr::MethodCall { receiver, args, .. } => {
524 Self::expr_uses_var(receiver, var_name)
525 || args.iter().any(|a| Self::expr_uses_var(a, var_name))
526 }
527 PureExpr::Call { func, args } => {
528 Self::expr_uses_var(func, var_name)
529 || args.iter().any(|a| Self::expr_uses_var(a, var_name))
530 }
531 PureExpr::Binary { left, right, .. } => {
532 Self::expr_uses_var(left, var_name) || Self::expr_uses_var(right, var_name)
533 }
534 PureExpr::If {
535 cond,
536 then_branch,
537 else_branch,
538 } => {
539 Self::expr_uses_var(cond, var_name)
540 || then_branch
541 .stmts
542 .iter()
543 .any(|s| Self::stmt_uses_var(s, var_name))
544 || else_branch
545 .as_ref()
546 .map(|e| Self::expr_uses_var(e, var_name))
547 .unwrap_or(false)
548 }
549 PureExpr::For { .. } => false, _ => false,
551 }
552 }
553
554 fn block_pattern_to_let_stmt(pattern: &BlockLoopPattern) -> PureStmt {
556 let iter_expr = Self::pattern_to_iter_expr(&pattern.loop_pattern);
557
558 PureStmt::Local {
559 pattern: PurePattern::Ident {
560 name: pattern.target_var.clone(),
561 is_mut: false, by_ref: false,
563 },
564 ty: Some(PureType::Other("Vec<_>".to_string())), init: Some(iter_expr),
566 else_branch: None,
567 }
568 }
569
570 pub fn transform_block(&self, block: &mut PureBlock) -> usize {
572 let mut changes = 0;
573
574 let block_patterns = Self::detect_block_patterns(&block.stmts);
576
577 if !block_patterns.is_empty() {
578 let mut indices_to_remove = Vec::new();
580
581 for pattern in block_patterns.iter().rev() {
582 let new_stmt = Self::block_pattern_to_let_stmt(pattern);
584 block.stmts[pattern.vec_decl_idx] = new_stmt;
585
586 indices_to_remove.push(pattern.for_loop_idx);
588 changes += 1;
589 }
590
591 indices_to_remove.sort();
593 indices_to_remove.reverse();
594 for idx in indices_to_remove {
595 block.stmts.remove(idx);
596 }
597 }
598
599 for stmt in &mut block.stmts {
601 changes += self.transform_stmt(stmt);
602 }
603
604 changes
605 }
606
607 fn transform_stmt(&self, stmt: &mut PureStmt) -> usize {
609 match stmt {
610 PureStmt::Semi(expr) | PureStmt::Expr(expr) => self.transform_expr(expr),
611 PureStmt::Local { init: Some(e), .. } => self.transform_expr(e),
612 _ => 0,
613 }
614 }
615
616 fn transform_expr(&self, expr: &mut PureExpr) -> usize {
618 let mut changes = 0;
619
620 match expr {
621 PureExpr::For {
622 pat,
623 expr: iter_expr,
624 body,
625 ..
626 } => {
627 if let Some(ref target) = self.target_var {
629 if let PurePattern::Ident { name, .. } = pat {
630 if name != target {
631 changes += self.transform_block(body);
633 return changes;
634 }
635 }
636 }
637
638 if self.aggressive {
640 if let Some(pattern) = Self::detect_pattern(pat, iter_expr, body) {
641 if matches!(pattern, LoopPattern::ForEach { .. }) {
642 *expr = Self::pattern_to_iter_expr(&pattern);
643 changes += 1;
644 }
645 }
646 } else {
647 changes += self.transform_block(body);
649 }
650 }
651 PureExpr::Block { block, .. } => {
652 changes += self.transform_block(block);
653 }
654 PureExpr::If {
655 cond,
656 then_branch,
657 else_branch,
658 } => {
659 changes += self.transform_expr(cond);
660 changes += self.transform_block(then_branch);
661 if let Some(else_expr) = else_branch {
662 changes += self.transform_expr(else_expr);
663 }
664 }
665 PureExpr::Match { expr: e, arms } => {
666 changes += self.transform_expr(e);
667 for arm in arms {
668 changes += self.transform_expr(&mut arm.body);
669 }
670 }
671 PureExpr::Loop { body: block, .. } | PureExpr::While { body: block, .. } => {
672 changes += self.transform_block(block);
673 }
674 PureExpr::Closure { body, .. } => {
675 changes += self.transform_expr(body);
676 }
677 _ => {}
678 }
679
680 changes
681 }
682}
683
684impl Mutation for LoopToIteratorMutation {
685 fn describe(&self) -> String {
686 "Convert for loops to iterator chains".to_string()
687 }
688
689 fn mutation_type(&self) -> &'static str {
690 "LoopToIterator"
691 }
692
693 fn box_clone(&self) -> Box<dyn Mutation> {
694 Box::new(self.clone())
695 }
696}
697
698#[cfg(test)]
699mod tests {
700 use super::*;
701
702 #[test]
703 fn test_detect_map_pattern() {
704 let pat = PurePattern::Ident {
705 name: "x".to_string(),
706 is_mut: false,
707 by_ref: false,
708 };
709 let iter = PureExpr::Path("items".to_string());
710 let body = PureBlock {
711 stmts: vec![PureStmt::Semi(PureExpr::MethodCall {
712 receiver: Box::new(PureExpr::Path("result".to_string())),
713 method: "push".to_string(),
714 turbofish: None,
715 args: vec![PureExpr::Binary {
716 op: "*".to_string(),
717 left: Box::new(PureExpr::Path("x".to_string())),
718 right: Box::new(PureExpr::Lit("2".to_string())),
719 }],
720 })],
721 };
722
723 let pattern = LoopToIteratorMutation::detect_pattern(&pat, &iter, &body);
724 assert!(matches!(pattern, Some(LoopPattern::MapCollect { .. })));
725 }
726
727 #[test]
728 fn test_detect_filter_pattern() {
729 let pat = PurePattern::Ident {
730 name: "x".to_string(),
731 is_mut: false,
732 by_ref: false,
733 };
734 let iter = PureExpr::Path("items".to_string());
735 let body = PureBlock {
736 stmts: vec![PureStmt::Semi(PureExpr::If {
737 cond: Box::new(PureExpr::Binary {
738 op: ">".to_string(),
739 left: Box::new(PureExpr::Path("x".to_string())),
740 right: Box::new(PureExpr::Lit("0".to_string())),
741 }),
742 then_branch: PureBlock {
743 stmts: vec![PureStmt::Semi(PureExpr::MethodCall {
744 receiver: Box::new(PureExpr::Path("result".to_string())),
745 method: "push".to_string(),
746 turbofish: None,
747 args: vec![PureExpr::Path("x".to_string())],
748 })],
749 },
750 else_branch: None,
751 })],
752 };
753
754 let pattern = LoopToIteratorMutation::detect_pattern(&pat, &iter, &body);
755 assert!(matches!(pattern, Some(LoopPattern::FilterCollect { .. })));
756 }
757}