Coverage for src / monte_neo / utils / ast_utils.py: 100%

49 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-28 16:27 +0200

1"""AST utilities for genetic programming. 

2 

3Handles parsing, manipulation, and unparsing of indicator code strings. 

4""" 

5 

6from __future__ import annotations 

7 

8import ast 

9import random 

10from typing import Any 

11 

12from monte_neo.utils.logger import get_logger 

13 

14logger = get_logger(__name__) 

15 

16 

17class ExpressionCollector(ast.NodeVisitor): 

18 """Collects all expression nodes from an AST.""" 

19 

20 def __init__(self) -> None: 

21 self.nodes: list[ast.AST] = [] 

22 

23 def visit_BinOp(self, node: ast.BinOp) -> Any: 

24 self.nodes.append(node) 

25 self.generic_visit(node) 

26 

27 def visit_Call(self, node: ast.Call) -> Any: 

28 self.nodes.append(node) 

29 self.generic_visit(node) 

30 

31 def visit_Attribute(self, node: ast.Attribute) -> Any: 

32 # data['close'].rolling(...) -> 'rolling' attribute is part of Call usually 

33 # but let's collect it too if it's top level 

34 self.nodes.append(node) 

35 self.generic_visit(node) 

36 

37 def visit_Subscript(self, node: ast.Subscript) -> Any: 

38 # data['close'] 

39 self.nodes.append(node) 

40 self.generic_visit(node) 

41 

42 

43class CrossoverTransformer(ast.NodeTransformer): 

44 """Replaces a specific node in the AST with another.""" 

45 

46 def __init__(self, target_node: ast.AST, replacement_node: ast.AST) -> None: 

47 self.target_node = target_node 

48 self.replacement_node = replacement_node 

49 

50 def visit(self, node: ast.AST) -> ast.AST: 

51 if node is self.target_node: 

52 return self.replacement_node 

53 return super().visit(node) 

54 

55 

56def crossover_trees(code1: str, code2: str) -> str: 

57 """Perform crossover between two code strings using AST subtree swapping. 

58 

59 Returns: 

60 New code string derived from code1 with a subtree from code2. 

61 """ 

62 try: 

63 tree1 = ast.parse(code1) 

64 tree2 = ast.parse(code2) 

65 

66 collector1 = ExpressionCollector() 

67 collector1.visit(tree1) 

68 

69 collector2 = ExpressionCollector() 

70 collector2.visit(tree2) 

71 

72 if not collector1.nodes or not collector2.nodes: 

73 logger.warning( 

74 "Crossover failed: No valid nodes found in one or both trees." 

75 ) 

76 return code1 

77 

78 # Pick random crossover points 

79 target_node = random.choice(collector1.nodes) 

80 replacement_node = random.choice(collector2.nodes) 

81 

82 # Transform tree1 by replacing target_node with replacement_node 

83 transformer = CrossoverTransformer(target_node, replacement_node) 

84 new_tree = transformer.visit(tree1) 

85 

86 # Fix missing line numbers etc. 

87 ast.fix_missing_locations(new_tree) 

88 

89 # Generate code back 

90 return ast.unparse(new_tree).strip() 

91 

92 except Exception as e: 

93 logger.error(f"Error during AST crossover: {e}") 

94 return code1