code-towel 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- code_towel-1.0.0.dist-info/METADATA +722 -0
- code_towel-1.0.0.dist-info/RECORD +27 -0
- code_towel-1.0.0.dist-info/WHEEL +5 -0
- code_towel-1.0.0.dist-info/entry_points.txt +4 -0
- code_towel-1.0.0.dist-info/licenses/LICENSE +201 -0
- code_towel-1.0.0.dist-info/top_level.txt +1 -0
- towel/__init__.py +20 -0
- towel/cli.py +813 -0
- towel/unification/__init__.py +24 -0
- towel/unification/assignment_analyzer.py +403 -0
- towel/unification/ast_normalizer.py +223 -0
- towel/unification/ast_pretty_printer.py +133 -0
- towel/unification/binding_detector.py +381 -0
- towel/unification/block_signature.py +123 -0
- towel/unification/builtins.py +209 -0
- towel/unification/exceptions.py +113 -0
- towel/unification/extractor.py +884 -0
- towel/unification/models.py +190 -0
- towel/unification/nominal_unifier.py +375 -0
- towel/unification/orphan_detector.py +187 -0
- towel/unification/pipeline.py +524 -0
- towel/unification/project_layout.py +201 -0
- towel/unification/refactor_engine.py +3381 -0
- towel/unification/scope_analyzer.py +597 -0
- towel/unification/unifier.py +2128 -0
- towel/unification/visitor_utils.py +83 -0
- towel/unification/visitors.py +377 -0
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
# Copyright 2025 Eric Allen
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
"""
|
|
16
|
+
Unification-based code refactoring.
|
|
17
|
+
|
|
18
|
+
This module implements a principled approach to detecting and extracting
|
|
19
|
+
duplicate code using unification from automated theorem proving.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
from .refactor_engine import UnificationRefactorEngine
|
|
23
|
+
|
|
24
|
+
__all__ = ["UnificationRefactorEngine"]
|
|
@@ -0,0 +1,403 @@
|
|
|
1
|
+
# Copyright 2025 Eric Allen
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
"""
|
|
16
|
+
Assignment analyzer for distinguishing initial bindings from reassignments.
|
|
17
|
+
|
|
18
|
+
This module analyzes assignment statements within a function to classify them as either:
|
|
19
|
+
- Initial bindings: First assignment to a variable (creates the variable)
|
|
20
|
+
- Reassignments: Subsequent assignments to an already-bound variable
|
|
21
|
+
|
|
22
|
+
This is critical for safe code extraction:
|
|
23
|
+
- Extracting code with an initial binding is safe
|
|
24
|
+
- Extracting code with a reassignment WITHOUT the initial binding is unsafe
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
import ast
|
|
28
|
+
from typing import Dict, List, Set, Tuple, Union
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def analyze_assignments(func: Union[ast.FunctionDef, ast.AsyncFunctionDef]) -> Dict[int, bool]:
|
|
32
|
+
"""
|
|
33
|
+
Analyze assignments in a function to identify reassignments.
|
|
34
|
+
|
|
35
|
+
Args:
|
|
36
|
+
func: Function definition to analyze
|
|
37
|
+
|
|
38
|
+
Returns:
|
|
39
|
+
Dictionary mapping assignment node id to is_reassignment boolean
|
|
40
|
+
- True: This assignment is a reassignment (variable was bound earlier)
|
|
41
|
+
- False: This assignment is an initial binding (first assignment to variable)
|
|
42
|
+
|
|
43
|
+
Example:
|
|
44
|
+
def foo(x):
|
|
45
|
+
result = x * 2 # Initial binding: id -> False
|
|
46
|
+
result = result + 10 # Reassignment: id -> True
|
|
47
|
+
return result
|
|
48
|
+
"""
|
|
49
|
+
analyzer = AssignmentAnalyzer()
|
|
50
|
+
analyzer.visit(func)
|
|
51
|
+
return analyzer.reassignments
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
class AssignmentAnalyzer(ast.NodeVisitor):
|
|
55
|
+
"""
|
|
56
|
+
Visitor that analyzes assignments to determine which are reassignments.
|
|
57
|
+
|
|
58
|
+
Tracks which variables have been bound and marks assignments accordingly.
|
|
59
|
+
Handles scoping correctly for nested functions, comprehensions, etc.
|
|
60
|
+
"""
|
|
61
|
+
|
|
62
|
+
def __init__(self) -> None:
|
|
63
|
+
self.bound_vars: Set[str] = set()
|
|
64
|
+
self.reassignments: Dict[int, bool] = {} # node id -> is_reassignment
|
|
65
|
+
|
|
66
|
+
def visit_FunctionDef(self, node: ast.FunctionDef) -> None:
|
|
67
|
+
"""
|
|
68
|
+
Visit function definition.
|
|
69
|
+
|
|
70
|
+
For the top-level function being analyzed:
|
|
71
|
+
- Parameters are considered bound variables
|
|
72
|
+
- Visit the function body
|
|
73
|
+
|
|
74
|
+
For nested functions:
|
|
75
|
+
- Don't descend (they have their own scope)
|
|
76
|
+
"""
|
|
77
|
+
# If this is the first function we're visiting, analyze it
|
|
78
|
+
if not self.bound_vars:
|
|
79
|
+
_record_function_parameters(node.args, self.bound_vars)
|
|
80
|
+
_record_kwonly_parameters(node.args, self.bound_vars)
|
|
81
|
+
|
|
82
|
+
# Visit function body
|
|
83
|
+
for stmt in node.body:
|
|
84
|
+
self.visit(stmt)
|
|
85
|
+
# else: Don't descend into nested functions (different scope)
|
|
86
|
+
|
|
87
|
+
def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> None:
|
|
88
|
+
"""Visit async function definition (same as FunctionDef)."""
|
|
89
|
+
# Type ignore needed because mypy doesn't recognize structural compatibility
|
|
90
|
+
self.visit_FunctionDef(node) # type: ignore[arg-type]
|
|
91
|
+
|
|
92
|
+
def visit_Assign(self, node: ast.Assign) -> None:
|
|
93
|
+
"""
|
|
94
|
+
Visit assignment statement.
|
|
95
|
+
|
|
96
|
+
For each target being assigned:
|
|
97
|
+
- If already bound: mark as reassignment
|
|
98
|
+
- If not bound: mark as initial binding and add to bound_vars
|
|
99
|
+
"""
|
|
100
|
+
# Visit the RHS first (in case it has side effects on bound vars)
|
|
101
|
+
self.visit(node.value)
|
|
102
|
+
|
|
103
|
+
# Process each target
|
|
104
|
+
for target in node.targets:
|
|
105
|
+
if isinstance(target, ast.Name):
|
|
106
|
+
# Simple variable assignment
|
|
107
|
+
var_name = target.id
|
|
108
|
+
|
|
109
|
+
# Check if this variable is already bound
|
|
110
|
+
is_reassignment = var_name in self.bound_vars
|
|
111
|
+
|
|
112
|
+
# Record the classification
|
|
113
|
+
self.reassignments[id(node)] = is_reassignment
|
|
114
|
+
|
|
115
|
+
# Mark variable as bound for future assignments
|
|
116
|
+
self.bound_vars.add(var_name)
|
|
117
|
+
else:
|
|
118
|
+
# Complex target (tuple unpacking, subscript, attribute)
|
|
119
|
+
# Collect any Name nodes being assigned to
|
|
120
|
+
names = self._collect_assignment_names(target)
|
|
121
|
+
|
|
122
|
+
# Check if ANY of the names are reassignments
|
|
123
|
+
is_any_reassignment = any(name in self.bound_vars for name in names)
|
|
124
|
+
self.reassignments[id(node)] = is_any_reassignment
|
|
125
|
+
|
|
126
|
+
# Mark all names as bound
|
|
127
|
+
self.bound_vars.update(names)
|
|
128
|
+
|
|
129
|
+
def visit_AugAssign(self, node: ast.AugAssign) -> None:
|
|
130
|
+
"""
|
|
131
|
+
Visit augmented assignment (+=, -=, etc.).
|
|
132
|
+
|
|
133
|
+
Augmented assignments are ALWAYS reassignments because they read
|
|
134
|
+
the variable before writing it.
|
|
135
|
+
"""
|
|
136
|
+
# The target must already be bound (or it's a runtime error)
|
|
137
|
+
# Mark as reassignment
|
|
138
|
+
self.reassignments[id(node)] = True
|
|
139
|
+
|
|
140
|
+
# Visit children
|
|
141
|
+
self.visit(node.target)
|
|
142
|
+
self.visit(node.value)
|
|
143
|
+
|
|
144
|
+
def visit_For(self, node: ast.For) -> None:
|
|
145
|
+
"""
|
|
146
|
+
Visit for loop.
|
|
147
|
+
|
|
148
|
+
The loop variable is bound by the for statement.
|
|
149
|
+
"""
|
|
150
|
+
# Visit the iterable first
|
|
151
|
+
self.visit(node.iter)
|
|
152
|
+
|
|
153
|
+
# The loop target creates bindings
|
|
154
|
+
if isinstance(node.target, ast.Name):
|
|
155
|
+
var_name = node.target.id
|
|
156
|
+
# This is an initial binding (for loop creates the variable)
|
|
157
|
+
# Note: We don't add a reassignments entry here because
|
|
158
|
+
# for loop targets are handled specially
|
|
159
|
+
self.bound_vars.add(var_name)
|
|
160
|
+
else:
|
|
161
|
+
# Complex target (tuple unpacking)
|
|
162
|
+
names = self._collect_assignment_names(node.target)
|
|
163
|
+
self.bound_vars.update(names)
|
|
164
|
+
|
|
165
|
+
# Visit loop body
|
|
166
|
+
_visit_body_and_orelse(self, node)
|
|
167
|
+
|
|
168
|
+
def visit_With(self, node: ast.With) -> None:
|
|
169
|
+
"""
|
|
170
|
+
Visit with statement.
|
|
171
|
+
|
|
172
|
+
The 'as' clause creates bindings.
|
|
173
|
+
"""
|
|
174
|
+
# Visit context expressions
|
|
175
|
+
for item in node.items:
|
|
176
|
+
self.visit(item.context_expr)
|
|
177
|
+
|
|
178
|
+
# The 'as' clause creates a binding
|
|
179
|
+
if item.optional_vars:
|
|
180
|
+
if isinstance(item.optional_vars, ast.Name):
|
|
181
|
+
self.bound_vars.add(item.optional_vars.id)
|
|
182
|
+
else:
|
|
183
|
+
names = self._collect_assignment_names(item.optional_vars)
|
|
184
|
+
self.bound_vars.update(names)
|
|
185
|
+
|
|
186
|
+
# Visit body
|
|
187
|
+
for stmt in node.body:
|
|
188
|
+
self.visit(stmt)
|
|
189
|
+
|
|
190
|
+
def visit_comprehension(self, node: ast.comprehension) -> None:
|
|
191
|
+
"""
|
|
192
|
+
Visit comprehension (in list/dict/set comprehension or generator).
|
|
193
|
+
|
|
194
|
+
Don't descend - comprehensions have their own scope.
|
|
195
|
+
"""
|
|
196
|
+
# Don't analyze comprehension targets as they create their own scope
|
|
197
|
+
pass
|
|
198
|
+
|
|
199
|
+
def visit_ListComp(self, node: ast.ListComp) -> None:
|
|
200
|
+
"""Don't descend into list comprehensions (own scope)."""
|
|
201
|
+
pass
|
|
202
|
+
|
|
203
|
+
def visit_DictComp(self, node: ast.DictComp) -> None:
|
|
204
|
+
"""Don't descend into dict comprehensions (own scope)."""
|
|
205
|
+
pass
|
|
206
|
+
|
|
207
|
+
def visit_SetComp(self, node: ast.SetComp) -> None:
|
|
208
|
+
"""Don't descend into set comprehensions (own scope)."""
|
|
209
|
+
pass
|
|
210
|
+
|
|
211
|
+
def visit_GeneratorExp(self, node: ast.GeneratorExp) -> None:
|
|
212
|
+
"""Don't descend into generator expressions (own scope)."""
|
|
213
|
+
pass
|
|
214
|
+
|
|
215
|
+
def _collect_assignment_names(self, target: ast.AST) -> Set[str]:
|
|
216
|
+
"""
|
|
217
|
+
Collect all Name nodes being assigned to in a complex target.
|
|
218
|
+
|
|
219
|
+
Examples:
|
|
220
|
+
- (a, b) = ... -> {'a', 'b'}
|
|
221
|
+
- [x, y, z] = ... -> {'x', 'y', 'z'}
|
|
222
|
+
- obj.attr = ... -> set() # Not a variable binding
|
|
223
|
+
- lst[i] = ... -> set() # Not a variable binding
|
|
224
|
+
"""
|
|
225
|
+
names: Set[str] = set()
|
|
226
|
+
|
|
227
|
+
class NameCollector(ast.NodeVisitor):
|
|
228
|
+
def visit_Name(self, node: ast.Name) -> None:
|
|
229
|
+
if isinstance(node.ctx, ast.Store):
|
|
230
|
+
names.add(node.id)
|
|
231
|
+
|
|
232
|
+
collector = NameCollector()
|
|
233
|
+
collector.visit(target)
|
|
234
|
+
return names
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def has_reassignments_without_bindings(
|
|
238
|
+
func: Union[ast.FunctionDef, ast.AsyncFunctionDef],
|
|
239
|
+
block_nodes: List[ast.AST],
|
|
240
|
+
reassignments: Dict[int, bool],
|
|
241
|
+
) -> Tuple[bool, Set[str]]:
|
|
242
|
+
"""
|
|
243
|
+
Check if a code block contains reassignments without initial bindings.
|
|
244
|
+
|
|
245
|
+
This is the validation function for safe extraction. A block is unsafe
|
|
246
|
+
to extract if it contains a reassignment to a variable that was initially
|
|
247
|
+
bound outside the block.
|
|
248
|
+
|
|
249
|
+
Args:
|
|
250
|
+
func: The function containing the block
|
|
251
|
+
block_nodes: The block being considered for extraction
|
|
252
|
+
reassignments: Assignment classification from analyze_assignments()
|
|
253
|
+
|
|
254
|
+
Returns:
|
|
255
|
+
Tuple of (has_unsafe_reassignments, set of problematic variable names)
|
|
256
|
+
- has_unsafe_reassignments: True if block is unsafe to extract
|
|
257
|
+
- problematic variables: Names of variables with reassignments but no bindings in block
|
|
258
|
+
|
|
259
|
+
Example:
|
|
260
|
+
def foo(x):
|
|
261
|
+
result = x * 2 # Line 2: initial binding
|
|
262
|
+
if result > 10: # Block starts here (line 3)
|
|
263
|
+
return result
|
|
264
|
+
result = result + 10 # Line 5: reassignment
|
|
265
|
+
return result # Block ends here
|
|
266
|
+
|
|
267
|
+
If we try to extract lines 3-6:
|
|
268
|
+
- Returns (True, {'result'}) because 'result' is reassigned on line 5
|
|
269
|
+
but initially bound on line 2 (outside the block)
|
|
270
|
+
"""
|
|
271
|
+
bound_in_block, reassigned_in_block = _collect_block_binding_stats(block_nodes, reassignments)
|
|
272
|
+
|
|
273
|
+
# Find variables that are reassigned but not initially bound in the block
|
|
274
|
+
problematic_vars = reassigned_in_block - bound_in_block
|
|
275
|
+
|
|
276
|
+
# Relaxation: allow reassignments to names declared global/nonlocal in the enclosing function
|
|
277
|
+
declared_global: Set[str] = set()
|
|
278
|
+
declared_nonlocal: Set[str] = set()
|
|
279
|
+
|
|
280
|
+
for stmt in func.body:
|
|
281
|
+
if isinstance(stmt, ast.Global):
|
|
282
|
+
declared_global.update(stmt.names)
|
|
283
|
+
elif isinstance(stmt, ast.Nonlocal):
|
|
284
|
+
declared_nonlocal.update(stmt.names)
|
|
285
|
+
|
|
286
|
+
allowed = declared_global | declared_nonlocal
|
|
287
|
+
remaining = problematic_vars - allowed
|
|
288
|
+
|
|
289
|
+
return (len(remaining) > 0, remaining)
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
def _collect_bindings_and_reassignments(
|
|
293
|
+
node: ast.AST, reassignments: Dict[int, bool], bound_vars: Set[str], reassigned_vars: Set[str]
|
|
294
|
+
) -> None:
|
|
295
|
+
"""
|
|
296
|
+
Recursively collect variables bound and reassigned in a node.
|
|
297
|
+
|
|
298
|
+
Args:
|
|
299
|
+
node: AST node to analyze
|
|
300
|
+
reassignments: Assignment classification mapping
|
|
301
|
+
bound_vars: Set to add initially-bound variables to
|
|
302
|
+
reassigned_vars: Set to add reassigned variables to
|
|
303
|
+
"""
|
|
304
|
+
|
|
305
|
+
class BindingCollector(ast.NodeVisitor):
|
|
306
|
+
def visit_Assign(self, node: ast.Assign) -> None:
|
|
307
|
+
is_reassignment = reassignments.get(id(node), False)
|
|
308
|
+
|
|
309
|
+
for target in node.targets:
|
|
310
|
+
if isinstance(target, ast.Name):
|
|
311
|
+
var_name = target.id
|
|
312
|
+
if is_reassignment:
|
|
313
|
+
reassigned_vars.add(var_name)
|
|
314
|
+
else:
|
|
315
|
+
bound_vars.add(var_name)
|
|
316
|
+
|
|
317
|
+
self.generic_visit(node)
|
|
318
|
+
|
|
319
|
+
def visit_AugAssign(self, node: ast.AugAssign) -> None:
|
|
320
|
+
# Augmented assignments are always reassignments
|
|
321
|
+
_add_augassign_target(node.target, reassigned_vars)
|
|
322
|
+
self.generic_visit(node)
|
|
323
|
+
|
|
324
|
+
def visit_For(self, node: ast.For) -> None:
|
|
325
|
+
# For loop variables are initial bindings
|
|
326
|
+
if isinstance(node.target, ast.Name):
|
|
327
|
+
bound_vars.add(node.target.id)
|
|
328
|
+
else:
|
|
329
|
+
# Complex target
|
|
330
|
+
class NameCollector(ast.NodeVisitor):
|
|
331
|
+
def visit_Name(self, n: ast.Name) -> None:
|
|
332
|
+
if isinstance(n.ctx, ast.Store):
|
|
333
|
+
bound_vars.add(n.id)
|
|
334
|
+
|
|
335
|
+
collector = NameCollector()
|
|
336
|
+
collector.visit(node.target)
|
|
337
|
+
|
|
338
|
+
self.generic_visit(node)
|
|
339
|
+
|
|
340
|
+
def visit_With(self, node: ast.With) -> None:
|
|
341
|
+
# With statement 'as' clauses create bindings
|
|
342
|
+
for item in node.items:
|
|
343
|
+
if item.optional_vars:
|
|
344
|
+
if isinstance(item.optional_vars, ast.Name):
|
|
345
|
+
bound_vars.add(item.optional_vars.id)
|
|
346
|
+
self.generic_visit(node)
|
|
347
|
+
|
|
348
|
+
def visit_FunctionDef(self, node: ast.FunctionDef) -> None:
|
|
349
|
+
# Don't descend into nested functions
|
|
350
|
+
pass
|
|
351
|
+
|
|
352
|
+
def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> None:
|
|
353
|
+
# Don't descend into nested async functions
|
|
354
|
+
pass
|
|
355
|
+
|
|
356
|
+
collector = BindingCollector()
|
|
357
|
+
collector.visit(node)
|
|
358
|
+
|
|
359
|
+
|
|
360
|
+
def _collect_block_binding_stats(
|
|
361
|
+
block_nodes: List[ast.AST], reassignments: Dict[int, bool]
|
|
362
|
+
) -> Tuple[Set[str], Set[str]]:
|
|
363
|
+
"""Return (bound_in_block, reassigned_in_block) for the given nodes."""
|
|
364
|
+
bound_in_block: Set[str] = set()
|
|
365
|
+
reassigned_in_block: Set[str] = set()
|
|
366
|
+
for node in block_nodes:
|
|
367
|
+
_collect_bindings_and_reassignments(
|
|
368
|
+
node, reassignments, bound_in_block, reassigned_in_block
|
|
369
|
+
)
|
|
370
|
+
return bound_in_block, reassigned_in_block
|
|
371
|
+
|
|
372
|
+
|
|
373
|
+
def _record_function_parameters(args: ast.arguments, target: Set[str]) -> None:
|
|
374
|
+
"""Add positional, vararg, and kwarg parameters to ``target``."""
|
|
375
|
+
for arg in args.args:
|
|
376
|
+
target.add(arg.arg)
|
|
377
|
+
if args.vararg:
|
|
378
|
+
target.add(args.vararg.arg)
|
|
379
|
+
if args.kwarg:
|
|
380
|
+
target.add(args.kwarg.arg)
|
|
381
|
+
|
|
382
|
+
|
|
383
|
+
def _record_kwonly_parameters(args: ast.arguments, target: Set[str]) -> None:
|
|
384
|
+
"""Add positional-only and keyword-only parameters to ``target``."""
|
|
385
|
+
for arg in args.posonlyargs:
|
|
386
|
+
target.add(arg.arg)
|
|
387
|
+
for arg in args.kwonlyargs:
|
|
388
|
+
target.add(arg.arg)
|
|
389
|
+
|
|
390
|
+
|
|
391
|
+
def _add_augassign_target(target: ast.AST, reassigned_vars: Set[str]) -> None:
|
|
392
|
+
"""Track Name targets that appear on the LHS of an augmented assignment."""
|
|
393
|
+
if isinstance(target, ast.Name):
|
|
394
|
+
reassigned_vars.add(target.id)
|
|
395
|
+
|
|
396
|
+
|
|
397
|
+
def _visit_body_and_orelse( # pragma: no cover - exercised via AssignmentAnalyzer
|
|
398
|
+
visitor: ast.NodeVisitor, node: ast.AST
|
|
399
|
+
) -> None:
|
|
400
|
+
for stmt in getattr(node, "body", []):
|
|
401
|
+
visitor.visit(stmt)
|
|
402
|
+
for stmt in getattr(node, "orelse", []):
|
|
403
|
+
visitor.visit(stmt)
|
|
@@ -0,0 +1,223 @@
|
|
|
1
|
+
# Copyright 2025 Eric Allen
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
"""
|
|
16
|
+
AST Normalizer - Convert assignments to augmented assignments when appropriate.
|
|
17
|
+
|
|
18
|
+
This module provides a normalization pass that identifies assignments of the form:
|
|
19
|
+
x = x + y
|
|
20
|
+
and converts them to augmented assignments:
|
|
21
|
+
x += y
|
|
22
|
+
|
|
23
|
+
This helps the unification algorithm distinguish between fresh bindings and mutations.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
import ast
|
|
27
|
+
from typing import Set, cast
|
|
28
|
+
from .visitor_utils import make_defensive_generic_visit
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class AssignToAugAssignNormalizer(ast.NodeTransformer):
|
|
32
|
+
"""
|
|
33
|
+
Normalize assignments to augmented assignments when the LHS variable
|
|
34
|
+
appears on the RHS and is already in scope.
|
|
35
|
+
|
|
36
|
+
Example transformations:
|
|
37
|
+
x = x + 1 → x += 1
|
|
38
|
+
result = result + 10 → result += 10
|
|
39
|
+
output = output - 5 → output -= 5
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
generic_visit = make_defensive_generic_visit("AssignToAugAssignNormalizer") # type: ignore[assignment]
|
|
43
|
+
|
|
44
|
+
def __init__(self) -> None:
|
|
45
|
+
self.scopes: list[Set[str]] = [set()] # Stack of scopes
|
|
46
|
+
|
|
47
|
+
def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef:
|
|
48
|
+
"""Enter a new scope for function definitions."""
|
|
49
|
+
# Add function parameters to scope
|
|
50
|
+
new_scope = set()
|
|
51
|
+
for arg in node.args.args:
|
|
52
|
+
new_scope.add(arg.arg)
|
|
53
|
+
|
|
54
|
+
self.scopes.append(new_scope)
|
|
55
|
+
|
|
56
|
+
# Visit the body
|
|
57
|
+
node.body = [self.visit(stmt) for stmt in node.body]
|
|
58
|
+
|
|
59
|
+
# Exit scope
|
|
60
|
+
self.scopes.pop()
|
|
61
|
+
|
|
62
|
+
return node
|
|
63
|
+
|
|
64
|
+
def visit_Assign(self, node: ast.Assign) -> ast.AST:
|
|
65
|
+
"""
|
|
66
|
+
Check if this assignment should be converted to an augmented assignment.
|
|
67
|
+
|
|
68
|
+
Conditions for conversion:
|
|
69
|
+
1. Single target (e.g., x = ...)
|
|
70
|
+
2. Target is a simple Name node (not subscript or attribute)
|
|
71
|
+
3. Target variable is already in scope
|
|
72
|
+
4. Value is a BinOp where one operand is the target variable
|
|
73
|
+
"""
|
|
74
|
+
# Visit children first
|
|
75
|
+
node = cast(ast.Assign, self.generic_visit(node))
|
|
76
|
+
|
|
77
|
+
# Check if we have a single target
|
|
78
|
+
if len(node.targets) != 1:
|
|
79
|
+
return node
|
|
80
|
+
|
|
81
|
+
target = node.targets[0]
|
|
82
|
+
|
|
83
|
+
# Check if target is a simple Name node
|
|
84
|
+
if not isinstance(target, ast.Name):
|
|
85
|
+
return node
|
|
86
|
+
|
|
87
|
+
target_name = target.id
|
|
88
|
+
|
|
89
|
+
# Check if target is already in scope
|
|
90
|
+
if not self._is_in_scope(target_name):
|
|
91
|
+
# This is a fresh binding, add to current scope
|
|
92
|
+
self.scopes[-1].add(target_name)
|
|
93
|
+
return node
|
|
94
|
+
|
|
95
|
+
# Target is in scope, check if value is a BinOp with target on LHS
|
|
96
|
+
if not isinstance(node.value, ast.BinOp):
|
|
97
|
+
return node
|
|
98
|
+
|
|
99
|
+
binop = node.value
|
|
100
|
+
|
|
101
|
+
# Check if left operand is the target variable
|
|
102
|
+
if isinstance(binop.left, ast.Name) and binop.left.id == target_name:
|
|
103
|
+
# Convert to augmented assignment: x = x + y → x += y
|
|
104
|
+
aug_assign = ast.AugAssign(
|
|
105
|
+
target=ast.Name(id=target_name, ctx=ast.Store()), op=binop.op, value=binop.right
|
|
106
|
+
)
|
|
107
|
+
return ast.copy_location(aug_assign, node)
|
|
108
|
+
|
|
109
|
+
# Check if right operand is the target variable (for commutative ops)
|
|
110
|
+
if (
|
|
111
|
+
self._is_commutative(binop.op)
|
|
112
|
+
and isinstance(binop.right, ast.Name)
|
|
113
|
+
and binop.right.id == target_name
|
|
114
|
+
):
|
|
115
|
+
# Convert: x = y + x → x += y
|
|
116
|
+
aug_assign = ast.AugAssign(
|
|
117
|
+
target=ast.Name(id=target_name, ctx=ast.Store()), op=binop.op, value=binop.left
|
|
118
|
+
)
|
|
119
|
+
return ast.copy_location(aug_assign, node)
|
|
120
|
+
|
|
121
|
+
return node
|
|
122
|
+
|
|
123
|
+
def visit_AugAssign(self, node: ast.AugAssign) -> ast.AugAssign:
|
|
124
|
+
"""Visit augmented assignment and track the variable."""
|
|
125
|
+
# The target of an augmented assignment must already be in scope
|
|
126
|
+
# (Python will raise NameError if it's not)
|
|
127
|
+
if isinstance(node.target, ast.Name):
|
|
128
|
+
# Ensure it's in scope (should already be, but add it if not)
|
|
129
|
+
self.scopes[-1].add(node.target.id)
|
|
130
|
+
|
|
131
|
+
return cast(ast.AugAssign, self.generic_visit(node))
|
|
132
|
+
|
|
133
|
+
def visit_For(self, node: ast.For) -> ast.For:
|
|
134
|
+
"""Track loop variables."""
|
|
135
|
+
if isinstance(node.target, ast.Name):
|
|
136
|
+
self.scopes[-1].add(node.target.id)
|
|
137
|
+
return cast(ast.For, self.generic_visit(node))
|
|
138
|
+
|
|
139
|
+
def visit_With(self, node: ast.With) -> ast.With:
|
|
140
|
+
"""Track context manager variables."""
|
|
141
|
+
for item in node.items:
|
|
142
|
+
if item.optional_vars and isinstance(item.optional_vars, ast.Name):
|
|
143
|
+
self.scopes[-1].add(item.optional_vars.id)
|
|
144
|
+
return cast(ast.With, self.generic_visit(node))
|
|
145
|
+
|
|
146
|
+
def visit_comprehension(self, node: ast.comprehension) -> ast.comprehension:
|
|
147
|
+
"""Track comprehension variables."""
|
|
148
|
+
if isinstance(node.target, ast.Name):
|
|
149
|
+
self.scopes[-1].add(node.target.id)
|
|
150
|
+
return cast(ast.comprehension, self.generic_visit(node))
|
|
151
|
+
|
|
152
|
+
def _is_in_scope(self, name: str) -> bool:
|
|
153
|
+
"""Check if a variable is in any of the current scopes."""
|
|
154
|
+
for scope in self.scopes:
|
|
155
|
+
if name in scope:
|
|
156
|
+
return True
|
|
157
|
+
return False
|
|
158
|
+
|
|
159
|
+
def _is_commutative(self, op: ast.operator) -> bool:
|
|
160
|
+
"""Check if an operator is commutative."""
|
|
161
|
+
return isinstance(op, (ast.Add, ast.Mult, ast.BitOr, ast.BitXor, ast.BitAnd))
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def normalize_assigns_to_augassigns(tree: ast.AST) -> ast.AST:
|
|
165
|
+
"""
|
|
166
|
+
Normalize an AST by converting assignments to augmented assignments
|
|
167
|
+
where appropriate.
|
|
168
|
+
|
|
169
|
+
Args:
|
|
170
|
+
tree: The AST to normalize
|
|
171
|
+
|
|
172
|
+
Returns:
|
|
173
|
+
A new AST with assignments normalized to augmented assignments
|
|
174
|
+
"""
|
|
175
|
+
normalizer = AssignToAugAssignNormalizer()
|
|
176
|
+
return cast(ast.AST, normalizer.visit(tree))
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
class ArithmeticCanonicalizer(ast.NodeTransformer):
|
|
180
|
+
"""Canonicalize arithmetic expressions for easier unification."""
|
|
181
|
+
|
|
182
|
+
generic_visit = make_defensive_generic_visit("ArithmeticCanonicalizer") # type: ignore[assignment]
|
|
183
|
+
|
|
184
|
+
def visit_UnaryOp(self, node: ast.UnaryOp) -> ast.AST: # noqa: N802
|
|
185
|
+
node = cast(ast.UnaryOp, self.generic_visit(node))
|
|
186
|
+
if isinstance(node.op, ast.USub) and isinstance(node.operand, ast.Constant):
|
|
187
|
+
value = node.operand.value
|
|
188
|
+
if isinstance(value, (int, float, complex)):
|
|
189
|
+
return ast.copy_location(ast.Constant(value=-value), node)
|
|
190
|
+
return node
|
|
191
|
+
|
|
192
|
+
def visit_BinOp(self, node: ast.BinOp) -> ast.AST: # noqa: N802
|
|
193
|
+
node = cast(ast.BinOp, self.generic_visit(node))
|
|
194
|
+
if isinstance(node.op, ast.Sub) and isinstance(node.right, ast.Constant):
|
|
195
|
+
value = node.right.value
|
|
196
|
+
if isinstance(value, (int, float, complex)):
|
|
197
|
+
new_const = ast.Constant(value=-value)
|
|
198
|
+
new_node = ast.BinOp(left=node.left, op=ast.Add(), right=new_const)
|
|
199
|
+
return ast.copy_location(new_node, node)
|
|
200
|
+
return node
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
def canonicalize_arithmetic(tree: ast.AST) -> ast.AST:
|
|
204
|
+
"""Apply arithmetic canonicalization for additive/subtractive expressions."""
|
|
205
|
+
|
|
206
|
+
canon = ArithmeticCanonicalizer()
|
|
207
|
+
return cast(ast.AST, canon.visit(tree))
|
|
208
|
+
|
|
209
|
+
|
|
210
|
+
def normalize_code(code: str) -> str:
|
|
211
|
+
"""
|
|
212
|
+
Normalize Python code by converting assignments to augmented assignments.
|
|
213
|
+
|
|
214
|
+
Args:
|
|
215
|
+
code: Python source code
|
|
216
|
+
|
|
217
|
+
Returns:
|
|
218
|
+
Normalized Python source code
|
|
219
|
+
"""
|
|
220
|
+
tree = ast.parse(code)
|
|
221
|
+
normalized_tree = normalize_assigns_to_augassigns(tree)
|
|
222
|
+
normalized_tree = canonicalize_arithmetic(normalized_tree)
|
|
223
|
+
return ast.unparse(normalized_tree)
|