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,884 @@
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
+ Hygienic code extraction.
17
+
18
+ Generates extracted functions while ensuring:
19
+ - No shadowing of enclosing scope identifiers
20
+ - Proper handling of hygienically renamed identifiers
21
+ - Preservation of evaluation order
22
+ - Referential transparency
23
+ """
24
+
25
+ import ast
26
+ import copy
27
+ from typing import List, Dict, Set, Tuple, Optional, TYPE_CHECKING, Callable, cast
28
+ from .unifier import Substitution
29
+
30
+ if TYPE_CHECKING:
31
+ from .scope_analyzer import Scope
32
+
33
+
34
+ class HygienicExtractor:
35
+ """
36
+ Extract code into a function while maintaining hygiene and
37
+ referential transparency.
38
+ """
39
+
40
+ def __init__(self) -> None:
41
+ self.used_names: Set[str] = set()
42
+
43
+ def extract_function(
44
+ self,
45
+ template_block: List[ast.AST],
46
+ substitution: Substitution,
47
+ free_variables: Set[str],
48
+ enclosing_names: Set[str],
49
+ is_value_producing: bool,
50
+ return_variables: Optional[List[str]] = None,
51
+ *,
52
+ global_decls: Optional[Set[str]] = None,
53
+ nonlocal_decls: Optional[Set[str]] = None,
54
+ function_name: str = "extracted_function",
55
+ ) -> Tuple[ast.FunctionDef, Dict[str, int]]:
56
+ """
57
+ Extract code into a function.
58
+
59
+ Args:
60
+ template_block: The code block to extract (from one of the blocks)
61
+ substitution: Substitution mapping expressions to parameters
62
+ free_variables: Free variables in the block
63
+ enclosing_names: Names defined in enclosing scopes
64
+ is_value_producing: Whether the block produces a value
65
+ return_variables: Variables to return from the extracted function (for value-producing extraction)
66
+ function_name: Name for the extracted function
67
+
68
+ Returns:
69
+ Tuple of (function AST node, parameter order dict)
70
+ """
71
+ # Reset name usage per extraction to keep function names stable across proposals
72
+ # and avoid cross-proposal suffix inflation.
73
+ self.used_names.clear()
74
+
75
+ if return_variables is None:
76
+ return_variables = []
77
+ # Ensure function name doesn't shadow
78
+ # Force double-underscore prefix for hygiene (avoid collisions with user code).
79
+ # If caller specified a different name explicitly, respect it; otherwise use the default.
80
+ if function_name == "__extracted_func":
81
+ function_name = self._ensure_unique_name(function_name, enclosing_names)
82
+ else:
83
+ # Still ensure uniqueness if a custom name was provided.
84
+ function_name = self._ensure_unique_name(function_name, enclosing_names)
85
+
86
+ # Determine parameters
87
+ # 1. Parameters from unification (substituted expressions)
88
+ # 2. Free variables (referenced but not bound in block)
89
+ # IMPORTANT: Keep unified parameter names EXACT (e.g., '__param_0') to remain
90
+ # consistent with Substitution lookups and replacements. Renaming these would
91
+ # desynchronize the body substitutions from the function signature.
92
+ param_names_unified = list(substitution.param_expressions.keys())
93
+
94
+ # Mapping from renamed to original names (identity since we don't rename)
95
+ rename_mapping = {name: name for name in param_names_unified}
96
+
97
+ # Add free variables as parameters (they're already unique)
98
+ param_names_free = sorted(free_variables)
99
+
100
+ # Combine: unified parameters first (to preserve evaluation order),
101
+ # then free variables
102
+ all_param_names = param_names_unified + param_names_free
103
+
104
+ # Create parameter order mapping
105
+ param_order = {name: idx for idx, name in enumerate(all_param_names)}
106
+
107
+ # Create function body by substituting unified parameters
108
+ body_nodes = self._substitute_parameters(
109
+ copy.deepcopy(template_block), substitution, param_names_unified, rename_mapping
110
+ )
111
+ # Substitute parameters returns generic AST nodes; for function body we expect statements
112
+ body: List[ast.stmt] = [cast(ast.stmt, n) for n in body_nodes]
113
+
114
+ # Detect parameters used as callees (in Call.func position) in the extracted body
115
+ # so we can safely defer their evaluation at call sites via zero-arg lambdas.
116
+ if param_names_unified:
117
+
118
+ class _CalleeParamFinder(ast.NodeVisitor):
119
+ def __init__(self, params: Set[str]) -> None:
120
+ self.params = params
121
+ self.found: Set[str] = set()
122
+
123
+ def visit_Call(self, node: ast.Call) -> None:
124
+ # If the callee is a Name matching a unified parameter, record it
125
+ if isinstance(node.func, ast.Name) and node.func.id in self.params:
126
+ self.found.add(node.func.id)
127
+ # Continue traversal
128
+ self.generic_visit(node)
129
+
130
+ finder = _CalleeParamFinder(set(param_names_unified))
131
+ for stmt in body:
132
+ finder.visit(stmt)
133
+ # Record on substitution for use during call generation
134
+ if hasattr(substitution, "params_used_as_callee"):
135
+ substitution.params_used_as_callee.update(finder.found)
136
+
137
+ # Optionally inject global/nonlocal declarations at the top of the extracted function
138
+ injected_preamble: List[ast.stmt] = []
139
+ if global_decls:
140
+ injected_preamble.append(ast.Global(names=sorted(global_decls)))
141
+ if nonlocal_decls:
142
+ injected_preamble.append(ast.Nonlocal(names=sorted(nonlocal_decls)))
143
+
144
+ # Add return statement for value-producing extraction
145
+ if return_variables:
146
+ # Prepare the return expression (expr type)
147
+ return_value: ast.expr
148
+ if len(return_variables) == 1:
149
+ # Single return variable: return var
150
+ return_value = ast.Name(id=return_variables[0], ctx=ast.Load())
151
+ else:
152
+ # Multiple return variables: return (var1, var2, ...)
153
+ return_value = ast.Tuple(
154
+ elts=[ast.Name(id=var, ctx=ast.Load()) for var in return_variables],
155
+ ctx=ast.Load(),
156
+ )
157
+
158
+ return_stmt = ast.Return(value=return_value)
159
+ body.append(return_stmt)
160
+
161
+ # Create function arguments
162
+ args = ast.arguments(
163
+ posonlyargs=[],
164
+ args=[ast.arg(arg=name) for name in all_param_names],
165
+ kwonlyargs=[],
166
+ kw_defaults=[],
167
+ defaults=[],
168
+ )
169
+
170
+ # Create function definition
171
+ # Prepend any injected declarations before the transformed body
172
+ final_body: List[ast.stmt] = (injected_preamble + body) if injected_preamble else body
173
+
174
+ func_def = ast.FunctionDef(
175
+ name=function_name,
176
+ args=args,
177
+ body=final_body if final_body else [ast.Pass()],
178
+ decorator_list=[],
179
+ returns=None,
180
+ )
181
+
182
+ # Fix missing locations
183
+ ast.fix_missing_locations(func_def)
184
+
185
+ return func_def, param_order
186
+
187
+ def generate_call(
188
+ self,
189
+ function_name: str,
190
+ block_idx: int,
191
+ substitution: Substitution,
192
+ param_order: Dict[str, int],
193
+ free_variables: Set[str],
194
+ is_value_producing: bool,
195
+ return_variables: Optional[List[str]] = None,
196
+ hygienic_renames: Optional[List[Dict[str, str]]] = None,
197
+ ) -> ast.stmt:
198
+ """
199
+ Generate a call to the extracted function.
200
+
201
+ Args:
202
+ function_name: Name of the function to call
203
+ block_idx: Index of the block being replaced
204
+ substitution: Substitution mapping
205
+ param_order: Parameter order from extract_function
206
+ free_variables: Free variables
207
+ is_value_producing: Whether this produces a value
208
+ return_variables: Variables that the extracted function returns
209
+ hygienic_renames: Hygienic renaming mapping for each block (original → canonical)
210
+
211
+ Returns:
212
+ AST node representing the call (either Return, Assign, or Expr)
213
+ """
214
+ if return_variables is None:
215
+ return_variables = []
216
+ if hygienic_renames is None or not hygienic_renames:
217
+ # Fallback: if the substitution carries hygienic renames, use them
218
+ if hasattr(substitution, "hygienic_renames") and substitution.hygienic_renames:
219
+ hygienic_renames = substitution.hygienic_renames
220
+ else:
221
+ hygienic_renames = []
222
+
223
+ # Build inverse mapping: canonical name → original name for this block
224
+ # hygienic_renames[block_idx] maps original → canonical, we need the reverse
225
+ inverse_renames: Dict[str, str] = {}
226
+ if block_idx < len(hygienic_renames):
227
+ for original_name, canonical_name in hygienic_renames[block_idx].items():
228
+ inverse_renames[canonical_name] = original_name
229
+
230
+ # Build arguments in correct order
231
+ # Build argument list (exprs); initialize as optional then cast when filled
232
+ args_list: List[Optional[ast.expr]] = [None] * len(param_order)
233
+
234
+ # Add unified parameters
235
+ for param_name, param_idx in param_order.items():
236
+ if param_name in substitution.param_expressions:
237
+ # This is a unified parameter - find the expression for this block
238
+ exprs = substitution.param_expressions[param_name]
239
+ for expr_block_idx, expr in exprs:
240
+ if expr_block_idx == block_idx:
241
+ # Check if this is a function parameter
242
+ if substitution.is_function_param(param_name):
243
+ # Wrap expression in lambda with bound variables
244
+ bound_vars = substitution.get_function_param_vars(param_name)
245
+ # Create lambda: lambda var1, var2, ...: expr
246
+ lambda_node = ast.Lambda(
247
+ args=ast.arguments(
248
+ posonlyargs=[],
249
+ args=[ast.arg(arg=var) for var in bound_vars],
250
+ kwonlyargs=[],
251
+ kw_defaults=[],
252
+ defaults=[],
253
+ ),
254
+ body=cast(ast.expr, expr),
255
+ )
256
+ args_list[param_idx] = lambda_node
257
+ elif (
258
+ hasattr(substitution, "params_used_as_callee")
259
+ and param_name in substitution.params_used_as_callee
260
+ ):
261
+ # Parameter is used as a callee in the extracted body (e.g., __param_0())
262
+ # Wrap it in a forwarding lambda that passes through any args/kwargs
263
+ # from the call site to the original callee expression.
264
+ # This avoids eager evaluation at the caller and preserves arity.
265
+ call_func = cast(ast.expr, expr if isinstance(expr, ast.expr) else expr)
266
+ call_body = ast.Call(
267
+ func=call_func,
268
+ args=[
269
+ ast.Starred(
270
+ value=ast.Name(id="args", ctx=ast.Load()), ctx=ast.Load()
271
+ )
272
+ ],
273
+ keywords=[
274
+ ast.keyword(
275
+ arg=None, value=ast.Name(id="kwargs", ctx=ast.Load())
276
+ )
277
+ ],
278
+ )
279
+
280
+ lambda_node = ast.Lambda(
281
+ args=ast.arguments(
282
+ posonlyargs=[],
283
+ args=[],
284
+ vararg=ast.arg(arg="args"),
285
+ kwonlyargs=[],
286
+ kw_defaults=[],
287
+ kwarg=ast.arg(arg="kwargs"),
288
+ defaults=[],
289
+ ),
290
+ body=call_body,
291
+ )
292
+ args_list[param_idx] = lambda_node
293
+ else:
294
+ # Regular parameter - use expression as-is
295
+ args_list[param_idx] = cast(ast.expr, expr)
296
+ break
297
+ else:
298
+ # This is a free variable - use the correct name for this block
299
+ # First check hygienic renames to find the original name for this block
300
+ var_name = inverse_renames.get(param_name, param_name)
301
+
302
+ # Also check if the name varies across blocks (augmented assignments)
303
+ if (
304
+ hasattr(substitution, "aug_assign_mappings")
305
+ and param_name in substitution.aug_assign_mappings
306
+ ):
307
+ mappings = substitution.aug_assign_mappings[param_name]
308
+ if block_idx in mappings:
309
+ var_name = mappings[block_idx]
310
+ args_list[param_idx] = ast.Name(id=var_name, ctx=ast.Load())
311
+ # If this parameter was introduced via higher-order literal promotion,
312
+ # prefer passing the original per-block expression rather than a free variable
313
+ # reference (which likely doesn't exist at the call site).
314
+ try:
315
+ if (
316
+ hasattr(substitution, "promoted_literal_args")
317
+ and substitution.promoted_literal_args
318
+ ):
319
+ promoted = substitution.promoted_literal_args.get(param_name, {})
320
+ if block_idx in promoted:
321
+ args_list[param_idx] = cast(ast.expr, promoted[block_idx])
322
+ except Exception:
323
+ # Best-effort; fall back to name reference on any issue
324
+ pass
325
+
326
+ # Create function call
327
+ call = ast.Call(
328
+ func=ast.Name(id=function_name, ctx=ast.Load()),
329
+ args=[cast(ast.expr, a) for a in args_list],
330
+ keywords=[],
331
+ )
332
+
333
+ # Map return variables to this block's original names when needed
334
+ mapped_return_vars: List[str] = []
335
+ if return_variables:
336
+ for var in return_variables:
337
+ # inverse_renames is Dict[str, str], default is the original var (str)
338
+ mapped_return_vars.append(inverse_renames.get(var, var))
339
+
340
+ # Handle wrapping based on return variables and is_value_producing
341
+ result_stmt: ast.stmt
342
+ if mapped_return_vars:
343
+ # Value-producing extraction with return variables
344
+ # Create assignment statement: result = func(args) or result, other = func(args)
345
+ if len(mapped_return_vars) == 1:
346
+ # Single variable: result = func(args)
347
+ assign_target: ast.expr = ast.Name(id=mapped_return_vars[0], ctx=ast.Store())
348
+ else:
349
+ # Multiple variables: result, other = func(args)
350
+ assign_target = ast.Tuple(
351
+ elts=[ast.Name(id=var, ctx=ast.Store()) for var in mapped_return_vars],
352
+ ctx=ast.Store(),
353
+ )
354
+ result_stmt = ast.Assign(targets=[assign_target], value=call)
355
+ elif is_value_producing:
356
+ # Value-producing extraction without return variables (has explicit return statements)
357
+ result_stmt = ast.Return(value=call)
358
+ else:
359
+ # Non-value-producing extraction
360
+ result_stmt = ast.Expr(value=call)
361
+
362
+ ast.fix_missing_locations(result_stmt)
363
+ return result_stmt
364
+
365
+ def _substitute_parameters(
366
+ self,
367
+ nodes: List[ast.AST],
368
+ substitution: Substitution,
369
+ param_names: List[str],
370
+ rename_mapping: Dict[str, str],
371
+ ) -> List[ast.AST]:
372
+ """
373
+ Substitute unified expressions with parameter names.
374
+
375
+ Args:
376
+ nodes: AST nodes to transform
377
+ substitution: Substitution mapping
378
+ param_names: Parameter names in order (renamed)
379
+ rename_mapping: Mapping from renamed to original parameter names
380
+
381
+ Returns:
382
+ Transformed AST nodes
383
+ """
384
+
385
+ # Create a transformer that replaces expressions with parameter names
386
+ class ParameterSubstituter(ast.NodeTransformer):
387
+ def __init__(
388
+ self, subst: Substitution, param_names: List[str], rename_mapping: Dict[str, str]
389
+ ) -> None:
390
+ self.subst = subst
391
+ self.param_names = param_names
392
+ self.rename_mapping = rename_mapping
393
+ # Use block 0 as the template
394
+ self.block_idx = 0
395
+ # Track if we're inside a JoinedStr to avoid breaking f-string structure
396
+ self.in_joinedstr = False
397
+ # Track variables that are equivalent to parameters
398
+ # Maps variable names to parameter names
399
+ self.var_to_param: Dict[str, str] = {}
400
+ # Track canonical parameter assigned to a variable name (even if shadowed later)
401
+ self.param_name_by_var: Dict[str, str] = {}
402
+ # Track parameterized variables that have been rebound to local values
403
+ self.shadowed_vars: Set[str] = set()
404
+
405
+ # CRITICAL: Initialize var_to_param with variables that are parameterized
406
+ # For each parameter, if its expression in block 0 is a simple variable name,
407
+ # then that variable should be substituted with the parameter throughout
408
+ for param_name in param_names:
409
+ # Get the original parameter name (before renaming)
410
+ original_param_name = rename_mapping.get(param_name, param_name)
411
+ if original_param_name in subst.param_expressions:
412
+ # This is a unified parameter - check if it's a simple variable reference
413
+ for block_idx, expr in subst.param_expressions[original_param_name]:
414
+ if block_idx == self.block_idx and isinstance(expr, ast.Name):
415
+ # This parameter represents a variable in our block
416
+ # Map the original variable name to the RENAMED parameter name
417
+ self.var_to_param[expr.id] = param_name
418
+ self.param_name_by_var[expr.id] = param_name
419
+ break
420
+
421
+ def _alias_variable(self, var_name: str, param_name: str) -> None:
422
+ self.var_to_param[var_name] = param_name
423
+ self.param_name_by_var[var_name] = param_name
424
+ self.shadowed_vars.discard(var_name)
425
+
426
+ def _mark_shadowed(self, var_name: str) -> None:
427
+ if var_name in self.param_name_by_var:
428
+ self.var_to_param.pop(var_name, None)
429
+ self.shadowed_vars.add(var_name)
430
+
431
+ def _variables_from_target(self, target: ast.AST) -> List[str]:
432
+ names: List[str] = []
433
+
434
+ def _collect(node: ast.AST) -> None:
435
+ if isinstance(node, ast.Name):
436
+ names.append(node.id)
437
+ elif isinstance(node, (ast.Tuple, ast.List)):
438
+ for elt in node.elts:
439
+ _collect(elt)
440
+
441
+ _collect(target)
442
+ return names
443
+
444
+ def _maybe_replace_node(self, node: ast.AST) -> Optional[ast.AST]:
445
+ # Only expressions participate in substitution mappings
446
+ if not isinstance(node, ast.expr):
447
+ return None
448
+
449
+ maybe_param_name: Optional[str] = self.subst.get_param_for_expr(
450
+ self.block_idx, node
451
+ )
452
+ if not maybe_param_name or maybe_param_name not in self.param_names:
453
+ return None
454
+
455
+ # CRITICAL: Never replace binding occurrences (Store/Del context)
456
+ if isinstance(node, ast.Name) and not isinstance(node.ctx, ast.Load):
457
+ return node
458
+
459
+ if isinstance(node, ast.Name) and node.id in self.shadowed_vars:
460
+ return node
461
+
462
+ # Inside an f-string literal component, keep constants intact
463
+ if self.in_joinedstr and isinstance(node, ast.Constant):
464
+ return node
465
+
466
+ # Don't replace FormattedValue nodes themselves; recurse into their value instead
467
+ if isinstance(node, ast.FormattedValue):
468
+ return None
469
+
470
+ if self.subst.is_function_param(maybe_param_name):
471
+ bound_vars = self.subst.get_function_param_vars(maybe_param_name)
472
+ call = ast.Call(
473
+ func=ast.Name(id=maybe_param_name, ctx=ast.Load()),
474
+ args=[ast.Name(id=var, ctx=ast.Load()) for var in bound_vars],
475
+ keywords=[],
476
+ )
477
+ return ast.copy_location(call, node)
478
+
479
+ # Regular parameter - just replace with parameter name
480
+ return ast.copy_location(ast.Name(id=maybe_param_name, ctx=ast.Load()), node)
481
+
482
+ def visit_JoinedStr(self, node: ast.JoinedStr) -> ast.JoinedStr:
483
+ # JoinedStr (f-string) can only have Constant or FormattedValue as direct children
484
+ # We must NEVER parameterize Constant nodes inside f-strings
485
+ # But we CAN parameterize expressions inside FormattedValue nodes
486
+ new_values: List[ast.expr] = []
487
+ previous_state = self.in_joinedstr
488
+ self.in_joinedstr = True
489
+ try:
490
+ for value in node.values:
491
+ if isinstance(value, ast.Constant):
492
+ # String literal parts of f-string must stay as constants
493
+ new_values.append(value)
494
+ elif isinstance(value, ast.FormattedValue):
495
+ # For FormattedValue, recursively visit the value expression
496
+ new_formatted = ast.FormattedValue(
497
+ value=cast(ast.expr, self.visit(value.value)),
498
+ conversion=value.conversion,
499
+ format_spec=value.format_spec,
500
+ )
501
+ new_values.append(new_formatted)
502
+ else:
503
+ # Shouldn't happen, but handle gracefully
504
+ new_values.append(cast(ast.expr, self.visit(value)))
505
+ finally:
506
+ self.in_joinedstr = previous_state
507
+ return ast.JoinedStr(values=new_values)
508
+
509
+ def visit_For(self, node: ast.For) -> ast.For:
510
+ """
511
+ Special handling for For loops to avoid replacing binding occurrences.
512
+
513
+ In 'for target in iter: body', the 'target' is a BINDING occurrence
514
+ and should NOT be replaced with a parameter.
515
+ """
516
+ # Transform the iterator (can contain parameterized expressions)
517
+ new_iter = cast(ast.expr, self.visit(node.iter))
518
+
519
+ # Don't transform the target (loop variable) - it's a binding
520
+ new_target = node.target
521
+ for var_name in self._variables_from_target(node.target):
522
+ self._mark_shadowed(var_name)
523
+
524
+ # Transform the body
525
+ new_body = self._visit_branch_statements(node.body)
526
+ new_orelse = self._visit_branch_statements(node.orelse) if node.orelse else []
527
+
528
+ return ast.For(target=new_target, iter=new_iter, body=new_body, orelse=new_orelse)
529
+
530
+ def visit_AsyncFor(self, node: ast.AsyncFor) -> ast.AsyncFor:
531
+ new_iter = cast(ast.expr, self.visit(node.iter))
532
+ new_target = node.target
533
+ for var_name in self._variables_from_target(node.target):
534
+ self._mark_shadowed(var_name)
535
+ new_body = self._visit_branch_statements(node.body)
536
+ new_orelse = self._visit_branch_statements(node.orelse) if node.orelse else []
537
+ return ast.AsyncFor(
538
+ target=new_target, iter=new_iter, body=new_body, orelse=new_orelse
539
+ )
540
+
541
+ def visit_comprehension(self, node: ast.comprehension) -> ast.comprehension:
542
+ """
543
+ Special handling for comprehensions to avoid replacing binding occurrences.
544
+
545
+ In 'for target in iter', the 'target' is a BINDING occurrence.
546
+ """
547
+ # Transform the iterator
548
+ new_iter = cast(ast.expr, self.visit(node.iter))
549
+
550
+ # Don't transform the target (comprehension variable) - it's a binding
551
+ new_target = node.target
552
+ for var_name in self._variables_from_target(node.target):
553
+ self._mark_shadowed(var_name)
554
+
555
+ # Transform the filters
556
+ new_ifs = [cast(ast.expr, self.visit(cond)) for cond in node.ifs]
557
+
558
+ return ast.comprehension(
559
+ target=new_target, iter=new_iter, ifs=new_ifs, is_async=node.is_async
560
+ )
561
+
562
+ def visit_Assign(self, node: ast.Assign) -> ast.Assign:
563
+ """
564
+ Special handling for assignments to handle both new bindings and reassignments.
565
+
566
+ For 'var = expr':
567
+ - If var is being assigned to a parameter (var = __param_N), track this mapping
568
+ - If var is in var_to_param and being reassigned to the SAME parameter, substitute target
569
+ - If var is in var_to_param but being reassigned to a DIFFERENT value, keep target as-is
570
+ and clear its mapping (creates new binding that shadows the parameter)
571
+ - Otherwise, keep the target unchanged (new binding)
572
+ """
573
+ # Transform the value expression first
574
+ new_value = cast(ast.expr, self.visit(node.value))
575
+
576
+ # Transform targets while preserving binding semantics
577
+ new_targets: List[ast.expr] = []
578
+ for target in node.targets:
579
+ if isinstance(target, ast.Name):
580
+ new_targets.append(target)
581
+ else:
582
+ new_targets.append(self._transform_assignment_target(target))
583
+
584
+ assigns_param = isinstance(new_value, ast.Name) and new_value.id in self.param_names
585
+ for target in node.targets:
586
+ for var_name in self._variables_from_target(target):
587
+ if assigns_param:
588
+ self._alias_variable(var_name, cast(ast.Name, new_value).id)
589
+ else:
590
+ self._mark_shadowed(var_name)
591
+
592
+ return ast.Assign(targets=new_targets, value=new_value)
593
+
594
+ def visit_If(self, node: ast.If) -> ast.If:
595
+ new_test = cast(ast.expr, self.visit(node.test))
596
+ new_body = self._visit_branch_statements(node.body)
597
+ new_orelse = self._visit_branch_statements(node.orelse)
598
+ return ast.If(test=new_test, body=new_body, orelse=new_orelse)
599
+
600
+ def visit_AugAssign(self, node: ast.AugAssign) -> ast.AugAssign:
601
+ new_value = cast(ast.expr, self.visit(node.value))
602
+ if isinstance(node.target, ast.Name):
603
+ self._mark_shadowed(node.target.id)
604
+ new_target: ast.Name | ast.Attribute | ast.Subscript = node.target
605
+ else:
606
+ new_target = cast(
607
+ ast.Attribute | ast.Subscript,
608
+ self._transform_assignment_target(node.target),
609
+ )
610
+ return ast.AugAssign(target=new_target, op=node.op, value=new_value)
611
+
612
+ def visit_With(self, node: ast.With) -> ast.With:
613
+ new_items = [
614
+ ast.withitem(
615
+ context_expr=cast(ast.expr, self.visit(item.context_expr)),
616
+ optional_vars=item.optional_vars,
617
+ )
618
+ for item in node.items
619
+ ]
620
+ new_body = self._visit_branch_statements(node.body)
621
+ return ast.With(items=new_items, body=new_body)
622
+
623
+ def visit_AsyncWith(self, node: ast.AsyncWith) -> ast.AsyncWith:
624
+ new_items = [
625
+ ast.withitem(
626
+ context_expr=cast(ast.expr, self.visit(item.context_expr)),
627
+ optional_vars=item.optional_vars,
628
+ )
629
+ for item in node.items
630
+ ]
631
+ new_body = self._visit_branch_statements(node.body)
632
+ return ast.AsyncWith(items=new_items, body=new_body)
633
+
634
+ def visit_While(self, node: ast.While) -> ast.While:
635
+ new_test = cast(ast.expr, self.visit(node.test))
636
+ new_body = self._visit_branch_statements(node.body)
637
+ new_orelse = self._visit_branch_statements(node.orelse) if node.orelse else []
638
+ return ast.While(test=new_test, body=new_body, orelse=new_orelse)
639
+
640
+ def visit_Try(self, node: ast.Try) -> ast.Try:
641
+ new_body = self._visit_branch_statements(node.body)
642
+ new_handlers = []
643
+ for handler in node.handlers:
644
+ new_type = cast(ast.expr, self.visit(handler.type)) if handler.type else None
645
+ new_handler_body = self._visit_branch_statements(handler.body)
646
+ new_handlers.append(
647
+ ast.ExceptHandler(type=new_type, name=handler.name, body=new_handler_body)
648
+ )
649
+ new_orelse = self._visit_branch_statements(node.orelse) if node.orelse else []
650
+ new_finalbody = (
651
+ self._visit_branch_statements(node.finalbody) if node.finalbody else []
652
+ )
653
+ return ast.Try(
654
+ body=new_body,
655
+ handlers=new_handlers,
656
+ orelse=new_orelse,
657
+ finalbody=new_finalbody,
658
+ )
659
+
660
+ def visit_AnnAssign(self, node: ast.AnnAssign) -> ast.AnnAssign:
661
+ new_value = cast(ast.expr, self.visit(node.value)) if node.value else None
662
+ if isinstance(node.target, (ast.Tuple, ast.List, ast.Attribute, ast.Subscript)):
663
+ new_target = self._transform_assignment_target(node.target)
664
+ else:
665
+ new_target = node.target
666
+ if isinstance(node.target, ast.Name):
667
+ if (
668
+ isinstance(new_value, ast.Name)
669
+ and new_value is not None
670
+ and new_value.id in self.param_names
671
+ ):
672
+ self._alias_variable(node.target.id, new_value.id)
673
+ else:
674
+ self._mark_shadowed(node.target.id)
675
+ return ast.AnnAssign(
676
+ target=new_target,
677
+ annotation=node.annotation,
678
+ value=new_value,
679
+ simple=node.simple,
680
+ )
681
+
682
+ def _transform_assignment_target(self, target: ast.expr) -> ast.expr:
683
+ """Recursively transform assignment targets while preserving binding semantics."""
684
+ if isinstance(target, ast.Name):
685
+ return target
686
+ if isinstance(target, (ast.Tuple, ast.List)):
687
+ new_elts = [self._transform_assignment_target(elt) for elt in target.elts]
688
+ return cast(
689
+ ast.expr,
690
+ ast.copy_location(type(target)(elts=new_elts, ctx=target.ctx), target),
691
+ )
692
+ if isinstance(target, ast.Attribute):
693
+ new_value = cast(ast.expr, self.visit(target.value))
694
+ return cast(
695
+ ast.expr,
696
+ ast.copy_location(
697
+ ast.Attribute(value=new_value, attr=target.attr, ctx=target.ctx), target
698
+ ),
699
+ )
700
+ if isinstance(target, ast.Subscript):
701
+ new_value = cast(ast.expr, self.visit(target.value))
702
+ new_slice = cast(ast.expr, self.visit(target.slice))
703
+ return cast(
704
+ ast.expr,
705
+ ast.copy_location(
706
+ ast.Subscript(value=new_value, slice=new_slice, ctx=target.ctx), target
707
+ ),
708
+ )
709
+ # Fallback: rely on generic_visit to transform child nodes
710
+ return cast(ast.expr, super().generic_visit(target))
711
+
712
+ def _visit_branch_statements(self, statements: List[ast.stmt]) -> List[ast.stmt]:
713
+ snapshot = self.var_to_param.copy()
714
+ shadow_snapshot = self.shadowed_vars.copy()
715
+ try:
716
+ result = [cast(ast.stmt, self.visit(stmt)) for stmt in statements]
717
+ current_state = self.var_to_param.copy()
718
+ finally:
719
+ current_state = locals().get("current_state", self.var_to_param.copy())
720
+ restored = snapshot.copy()
721
+ for var_name, param_name in list(snapshot.items()):
722
+ if var_name not in current_state:
723
+ restored.pop(var_name, None)
724
+ elif current_state[var_name] != param_name:
725
+ restored.pop(var_name, None)
726
+ self.var_to_param = restored
727
+ current_shadowed = self.shadowed_vars.copy()
728
+ self.shadowed_vars = shadow_snapshot | current_shadowed
729
+ return result
730
+
731
+ def visit(self, node: ast.AST) -> ast.AST:
732
+ replacement = self._maybe_replace_node(node)
733
+ if replacement is not None:
734
+ return replacement
735
+
736
+ method_name = f"visit_{node.__class__.__name__}"
737
+ visitor = getattr(self, method_name, None)
738
+ if visitor is None:
739
+ generic_result = super().generic_visit(node)
740
+ return generic_result
741
+ visit_callable = cast(Callable[[ast.AST], ast.AST], visitor)
742
+ return visit_callable(node)
743
+
744
+ substituter = ParameterSubstituter(substitution, param_names, rename_mapping)
745
+ return [substituter.visit(node) for node in nodes]
746
+
747
+ def _ensure_unique_name(self, name: str, enclosing_names: Set[str]) -> str:
748
+ """
749
+ Ensure a name doesn't shadow enclosing scope names.
750
+
751
+ Args:
752
+ name: Proposed name
753
+ enclosing_names: Names in enclosing scopes
754
+
755
+ Returns:
756
+ Unique name (possibly with numeric suffix)
757
+ """
758
+ if name not in enclosing_names and name not in self.used_names:
759
+ self.used_names.add(name)
760
+ return name
761
+
762
+ # Add numeric suffix with __ prefix to avoid name collisions
763
+ counter = 1
764
+ while True:
765
+ candidate = f"__{name}_{counter}"
766
+ if candidate not in enclosing_names and candidate not in self.used_names:
767
+ self.used_names.add(candidate)
768
+ return candidate
769
+ counter += 1
770
+
771
+
772
+ def contains_return(block: List[ast.stmt]) -> bool:
773
+ """
774
+ Check if a block contains any return statements (including nested ones).
775
+ """
776
+
777
+ class ReturnFinder(ast.NodeVisitor):
778
+ def __init__(self) -> None:
779
+ self.found_return: bool = False
780
+
781
+ def visit_Return(self, node: ast.Return) -> None:
782
+ self.found_return = True
783
+
784
+ def visit_FunctionDef(self, node: ast.FunctionDef) -> None:
785
+ # Don't visit nested function definitions
786
+ pass
787
+
788
+ def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> None:
789
+ # Don't visit nested async function definitions
790
+ pass
791
+
792
+ finder = ReturnFinder()
793
+ for stmt in block:
794
+ finder.visit(stmt)
795
+ if finder.found_return:
796
+ return True
797
+ return False
798
+
799
+
800
+ def is_value_producing(block: List[ast.stmt]) -> bool:
801
+ """
802
+ Check if a block of code produces a value.
803
+
804
+ A block is value-producing if:
805
+ - It contains a return statement (including nested)
806
+ - It's a single expression
807
+ """
808
+ if not block:
809
+ return False
810
+
811
+ # Check if block contains any return statements
812
+ if contains_return(block):
813
+ return True
814
+
815
+ # Single expression statement
816
+ last_stmt = block[-1]
817
+ if len(block) == 1 and isinstance(last_stmt, ast.Expr):
818
+ return True
819
+
820
+ return False
821
+
822
+
823
+ def has_complete_return_coverage(block: List[ast.stmt]) -> bool:
824
+ """
825
+ Check if a value-producing block has complete return coverage.
826
+
827
+ This ensures that if a block contains conditional returns (like an IF
828
+ with a return in the if-branch), it also has a return for the else case.
829
+
830
+ Returns True if:
831
+ - The last statement is a return, OR
832
+ - The last statement is an IF/While/For with returns in ALL branches
833
+
834
+ Args:
835
+ block: List of AST statements
836
+
837
+ Returns:
838
+ True if the block has complete return coverage
839
+ """
840
+ if not block:
841
+ return False
842
+
843
+ last_stmt = block[-1]
844
+
845
+ # If the last statement is a return, we have complete coverage
846
+ if isinstance(last_stmt, ast.Return):
847
+ return True
848
+
849
+ # If the last statement is an IF
850
+ if isinstance(last_stmt, ast.If):
851
+ # Check if both branches have returns
852
+ if_has_return = contains_return(last_stmt.body)
853
+
854
+ # Check else branch
855
+ if last_stmt.orelse:
856
+ else_has_return = contains_return(last_stmt.orelse)
857
+ # Complete coverage if both branches return
858
+ return if_has_return and else_has_return
859
+ else:
860
+ # No else branch - incomplete coverage unless there's a return after
861
+ return False
862
+
863
+ # For other control structures (while, for, etc), incomplete coverage
864
+ # unless there's a return after them
865
+ return False
866
+
867
+
868
+ def get_enclosing_names(scope_tree: "Scope", current_scope: "Scope") -> Set[str]:
869
+ """
870
+ Get all names defined in scopes enclosing the current scope.
871
+
872
+ Args:
873
+ scope_tree: Root scope
874
+ current_scope: Current scope
875
+
876
+ Returns:
877
+ Set of names in enclosing scopes
878
+ """
879
+ names: Set[str] = set()
880
+ scope = current_scope.parent
881
+ while scope is not None:
882
+ names.update(scope.bindings.keys())
883
+ scope = scope.parent
884
+ return names