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
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
1"""AST utilities for genetic programming.
3Handles parsing, manipulation, and unparsing of indicator code strings.
4"""
6from __future__ import annotations
8import ast
9import random
10from typing import Any
12from monte_neo.utils.logger import get_logger
14logger = get_logger(__name__)
17class ExpressionCollector(ast.NodeVisitor):
18 """Collects all expression nodes from an AST."""
20 def __init__(self) -> None:
21 self.nodes: list[ast.AST] = []
23 def visit_BinOp(self, node: ast.BinOp) -> Any:
24 self.nodes.append(node)
25 self.generic_visit(node)
27 def visit_Call(self, node: ast.Call) -> Any:
28 self.nodes.append(node)
29 self.generic_visit(node)
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)
37 def visit_Subscript(self, node: ast.Subscript) -> Any:
38 # data['close']
39 self.nodes.append(node)
40 self.generic_visit(node)
43class CrossoverTransformer(ast.NodeTransformer):
44 """Replaces a specific node in the AST with another."""
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
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)
56def crossover_trees(code1: str, code2: str) -> str:
57 """Perform crossover between two code strings using AST subtree swapping.
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)
66 collector1 = ExpressionCollector()
67 collector1.visit(tree1)
69 collector2 = ExpressionCollector()
70 collector2.visit(tree2)
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
78 # Pick random crossover points
79 target_node = random.choice(collector1.nodes)
80 replacement_node = random.choice(collector2.nodes)
82 # Transform tree1 by replacing target_node with replacement_node
83 transformer = CrossoverTransformer(target_node, replacement_node)
84 new_tree = transformer.visit(tree1)
86 # Fix missing line numbers etc.
87 ast.fix_missing_locations(new_tree)
89 # Generate code back
90 return ast.unparse(new_tree).strip()
92 except Exception as e:
93 logger.error(f"Error during AST crossover: {e}")
94 return code1