Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 73 additions & 25 deletions diffgraph/structural.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,7 @@ class _Symbol:
class _Import:
module: str
line: int
position: int
snippet: str
bindings: Tuple[str, ...]
scope: Optional[str]
Expand All @@ -91,6 +92,7 @@ class _Call:
caller: Optional[str]
name: str
line: int
position: int
snippet: str
comprehension_bindings: Tuple[str, ...] = ()

Expand Down Expand Up @@ -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"):
Expand Down Expand Up @@ -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":
Expand All @@ -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:
Expand Down Expand Up @@ -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 (
Expand Down Expand Up @@ -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),
)
Expand All @@ -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"
Comment thread
coderabbitai[bot] marked this conversation as resolved.
):
# 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"
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
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"
Expand All @@ -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)

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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])
Expand Down
107 changes: 107 additions & 0 deletions tests/test_structural.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
Loading