diff --git a/diffgraph/structural.py b/diffgraph/structural.py index fdfca51..dbff6b8 100644 --- a/diffgraph/structural.py +++ b/diffgraph/structural.py @@ -81,6 +81,7 @@ class _Symbol: class _Import: module: str line: int + position: int snippet: str bindings: Tuple[str, ...] scope: Optional[str] @@ -91,6 +92,7 @@ class _Call: caller: Optional[str] name: str line: int + position: int snippet: str comprehension_bindings: Tuple[str, ...] = () @@ -375,6 +377,28 @@ def enclosing_comprehension_bindings(node) -> Tuple[str, ...]: ancestor = ancestor.parent return tuple(sorted(found)) + def is_inside_lambda(node) -> bool: + """Return whether node belongs to a lambda's lexical scope.""" + ancestor = node.parent + while ancestor is not None: + if ancestor.type == "lambda": + return True + ancestor = ancestor.parent + return False + + def is_declaration_header(node) -> bool: + """Return whether node is evaluated before its declaration binds.""" + ancestor = node.parent + while ancestor is not None: + if ancestor.type not in ("class_definition", "function_definition"): + ancestor = ancestor.parent + continue + body = ancestor.child_by_field_name("body") + return body is None or not ( + body.start_byte <= node.start_byte and node.end_byte <= body.end_byte + ) + return False + def visit(node, parents: Tuple[Tuple[str, str], ...] = ()) -> None: next_parents = parents if node.type in ("class_definition", "function_definition"): @@ -425,8 +449,15 @@ def visit(node, parents: Tuple[Tuple[str, str], ...] = ()) -> None: next_parents = (*parents, (qname, kind)) if not parents: # A top-level declaration overwrites an imported binding at - # runtime just like a top-level assignment does. - module_rebindings.append((name, node.start_point[0] + 1)) + # runtime just like a top-level assignment does, but only + # after the declaration header has been evaluated. + body = node.child_by_field_name("body") + binding_position = ( + body.start_byte + if node.type == "function_definition" and body is not None + else node.end_byte + ) + module_rebindings.append((name, binding_position)) elif node.type in ("import_statement", "import_from_statement"): snippet = _node_text(content, node) if node.type == "import_statement": @@ -445,7 +476,7 @@ def visit(node, parents: Tuple[Tuple[str, str], ...] = ()) -> None: # ``import package.submodule`` binds ``package``. binding = raw.split(".", 1)[0] imports.append(_Import( - raw, node.start_point[0] + 1, snippet, (binding,), + raw, node.start_point[0] + 1, node.end_byte, snippet, (binding,), parents[-1][0] if parents else None, )) else: @@ -478,10 +509,8 @@ def visit(node, parents: Tuple[Tuple[str, str], ...] = ()) -> None: else: imported.append(_node_text(content, child).split(".", 1)[0]) imports.append(_Import( - _node_text(content, module_node), - node.start_point[0] + 1, - snippet, - tuple(imported), + _node_text(content, module_node), node.start_point[0] + 1, + node.end_byte, snippet, tuple(imported), parents[-1][0] if parents else None, )) elif not parents and node.type in ( @@ -547,6 +576,7 @@ def visit(node, parents: Tuple[Tuple[str, str], ...] = ()) -> None: caller, _node_text(content, function), node.start_point[0] + 1, + node.start_byte, _node_text(content, node), enclosing_comprehension_bindings(node), ) @@ -565,24 +595,44 @@ def visit(node, parents: Tuple[Tuple[str, str], ...] = ()) -> None: name = _node_text(content, name_node) bindings.setdefault(scope, set()).add(name) if scope is None: - module_rebindings.append((name, node.start_point[0] + 1)) + module_rebindings.append((name, node.start_byte)) elif node.type == "case_clause": # Capture patterns bind their names before the case body executes. # Keep them lexical and line-aware just like assignment bindings. bound_names = case_pattern_identifiers(node) bindings.setdefault(scope, set()).update(bound_names) if scope is None: - module_rebindings.extend( - (name, node.start_point[0] + 1) for name in bound_names - ) - elif node.type in ("assignment", "annotated_assignment", "for_statement"): - left = node.child_by_field_name("left") + module_rebindings.extend((name, node.start_byte) for name in bound_names) + elif node.type in ( + "assignment", "annotated_assignment", "for_statement", "named_expression" + ): + # Assignment expressions (``name := value``) bind their target in + # the enclosing scope before later expressions execute. Treat them + # like ordinary assignments so an imported name cannot produce a + # false import-grounded call edge after it has been rebound. + left = node.child_by_field_name( + "name" if node.type == "named_expression" else "left" + ) if left is not None: bound_names = identifiers(left) - bindings.setdefault(scope, set()).update(bound_names) - if scope is None: + binding_scope = scope + if node.type == "named_expression": + # Lambdas have their own lexical scope, but are not + # structural symbols. Do not let their local targets leak + # into the enclosing function or module. + if is_inside_lambda(node): + bound_names = set() + elif is_declaration_header(node): + binding_scope = parents[-2][0] if len(parents) > 1 else None + bindings.setdefault(binding_scope, set()).update(bound_names) + if binding_scope is None: + binding_position = node.end_byte + if node.type == "for_statement": + iterable = node.child_by_field_name("right") + if iterable is not None: + binding_position = iterable.end_byte module_rebindings.extend( - (name, node.start_point[0] + 1) for name in bound_names + (name, binding_position) for name in bound_names ) elif ( node.type == "as_pattern" @@ -601,9 +651,7 @@ def visit(node, parents: Tuple[Tuple[str, str], ...] = ()) -> None: bound_names = as_target_identifiers(target) bindings.setdefault(scope, set()).update(bound_names) if scope is None: - module_rebindings.extend( - (name, node.start_point[0] + 1) for name in bound_names - ) + module_rebindings.extend((name, node.start_byte) for name in bound_names) for child in node.children: visit(child, next_parents) @@ -714,7 +762,7 @@ def _resolve_call_target( if current.kind in ("function", "method") or is_initial_class_body: if call.name in bindings.get(current_name, set()): history = imported_targets.get(current_name, {}).get(call.name, []) - visible = [target for line, target in history if line <= call.line] + visible = [target for position, target in history if position <= call.position] if visible: return visible[-1] return None @@ -728,7 +776,7 @@ def _resolve_call_target( # Select the binding visible at this call site rather than applying a # later top-level rebind retroactively. history = imported_targets.get(None, {}).get(call.name, []) - visible = [target for line, target in history if line <= call.line] + visible = [target for position, target in history if position <= call.position] if visible and visible[-1] is not None: return visible[-1] # A later import must not hide a declaration that was already visible @@ -773,13 +821,13 @@ def _imported_call_targets( # A later import of the same local name is intentionally # unresolved, but calls before it retain the earlier binding. scope_targets.setdefault(binding, []).append(( - item.line, None if binding in scope_bindings else target + item.position, None if binding in scope_bindings else target )) scope_bindings.add(binding) - for binding, line in module_rebindings: + for binding, position in module_rebindings: # A declaration, assignment, or loop target replaces the imported - # binding only for calls at or after its source line. - targets.setdefault(None, {}).setdefault(binding, []).append((line, None)) + # binding after its source expression has been evaluated. + targets.setdefault(None, {}).setdefault(binding, []).append((position, None)) for scope_targets in targets.values(): for history in scope_targets.values(): history.sort(key=lambda item: item[0]) diff --git a/tests/test_structural.py b/tests/test_structural.py index d89b232..2e47aec 100644 --- a/tests/test_structural.py +++ b/tests/test_structural.py @@ -1492,6 +1492,70 @@ def test_as_pattern_bindings_do_not_create_import_grounded_call_edges(tmp_path): assert calls == [] +def test_named_expression_bindings_do_not_create_import_grounded_call_edges(tmp_path): + """A walrus target shadows an import for calls after the assignment.""" + root = repo(tmp_path) + write( + root, + "named_expression_bindings.py", + "from remote.worker import execute as run_remote\n\n" + "def factory():\n" + " return lambda: None\n\n" + "def caller():\n" + " if (run_remote := factory()):\n" + " run_remote()\n", + ) + git(root, "add", "named_expression_bindings.py") + + artifact = analyze_local_diff(str(root), staged=True) + + assert_valid(artifact) + calls = [item for item in artifact["relationships"] if item["kind"] == "calls"] + assert [(item["source_id"], item["target_id"]) for item in calls] == [ + ("sym::named_expression_bindings.py::caller", "sym::named_expression_bindings.py::factory") + ] + + +def test_module_named_expression_keeps_import_visible_during_its_value(tmp_path): + """A walrus rebind takes effect after its value expression is evaluated.""" + root = repo(tmp_path) + write( + root, + "named_expression_order.py", + "from remote.worker import execute as run_remote\n\n" + "if (run_remote := run_remote()):\n" + " pass\n", + ) + git(root, "add", "named_expression_order.py") + + artifact = analyze_local_diff(str(root), staged=True) + + assert_valid(artifact) + calls = [item for item in artifact["relationships"] if item["kind"] == "calls"] + assert len(calls) == 1 + assert calls[0]["resolution_method"] == "import_grounded" + + +def test_module_for_target_keeps_import_visible_during_iterable_evaluation(tmp_path): + """A loop target binds after its iterable expression is evaluated.""" + root = repo(tmp_path) + write( + root, + "for_binding_order.py", + "from remote.worker import execute as run_remote\n\n" + "for run_remote in range(run_remote()):\n" + " pass\n", + ) + git(root, "add", "for_binding_order.py") + + artifact = analyze_local_diff(str(root), staged=True) + + assert_valid(artifact) + calls = [item for item in artifact["relationships"] if item["kind"] == "calls"] + assert len(calls) == 1 + assert calls[0]["resolution_method"] == "import_grounded" + + def test_comprehension_targets_shadow_imports_only_inside_comprehensions(tmp_path): """Comprehension targets shadow imports without leaking into their function.""" root = repo(tmp_path) @@ -1681,6 +1745,49 @@ def test_as_binding_destructuring_rebinds_each_imported_alias(tmp_path): assert calls == [] +def test_lambda_named_expression_does_not_shadow_an_enclosing_import(tmp_path): + """A lambda-local walrus target must not leak into the enclosing function.""" + root = repo(tmp_path) + write( + root, + "lambda_named_expression.py", + "from remote.worker import execute as run_remote\n\n" + "def caller():\n" + " thunk = lambda: (run_remote := 1)\n" + " return run_remote()\n", + ) + git(root, "add", "lambda_named_expression.py") + + artifact = analyze_local_diff(str(root), staged=True) + + assert_valid(artifact) + calls = [item for item in artifact["relationships"] if item["kind"] == "calls"] + assert len(calls) == 1 + assert calls[0]["resolution_method"] == "import_grounded" + + +def test_declaration_headers_keep_imports_visible_before_rebinding(tmp_path): + """Default and base expressions run before their top-level names bind.""" + root = repo(tmp_path) + write( + root, + "declaration_header_order.py", + "from remote.worker import execute as run_remote\n\n" + "def function_rebind(value=run_remote()):\n" + " pass\n\n" + "class class_rebind(run_remote()):\n" + " pass\n", + ) + git(root, "add", "declaration_header_order.py") + + artifact = analyze_local_diff(str(root), staged=True) + + assert_valid(artifact) + calls = [item for item in artifact["relationships"] if item["kind"] == "calls"] + assert len(calls) == 2 + assert all(item["resolution_method"] == "import_grounded" for item in calls) + + @pytest.mark.parametrize("declaration", ["def run_remote():\n return None", "class run_remote:\n pass"]) def test_module_declaration_rebinds_imported_alias(tmp_path, declaration): root = repo(tmp_path)