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.
@@ -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)