fix(extract): attribute cross-file INFERRED uses edges to the referencing symbol (#2652)

Rewrites Pass 2 of the Python cross-file import resolver so an INFERRED
`uses` edge anchors on the symbol whose body actually references the
imported name (a class as a unit, or a module-level function) at the real
reference line, instead of fanning out from the import statement line to
every class in the importing file.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Ousama Ben Younes
2026-08-14 14:17:45 +01:00
committed by safishamsi
co-authored by Claude Opus 4.8
parent 0302bfa7af
commit c08d9afa93
2 changed files with 173 additions and 79 deletions
+113 -79
View File
@@ -1935,25 +1935,31 @@ def _resolve_cross_file_imports(
if src_path.stem not in bare_to_qualified:
bare_to_qualified[src_path.stem] = fq_stem
# Pass 2: for each file, find `from .X import A, B, C` and resolve
# Pass 2: for each file, find `from .X import A, B, C`, then attribute the
# `uses` edge to the specific local symbol (class OR function) whose body
# actually references the imported name — not to every class that merely
# shares the file (#2652). The edge is anchored at the real reference, not
# the import line, so `source_location` points at genuine corroboration.
new_edges: list[dict] = []
stem_to_path: dict[str, Path] = {_file_stem(p): p for p in paths}
for file_result, path in zip(per_file, paths):
stem = _file_stem(path)
str_path = str(path)
# Find all classes defined in this file (the importers).
# Excludes rationale nodes whose labels happen not to end in ")" or ".py"
# but which must never be treated as importing entities (#563).
local_classes = [
n["id"] for n in file_result.get("nodes", [])
if n.get("source_file") == str_path
and not n["label"].endswith((")", ".py"))
and n["id"] != _make_id(stem) # exclude file-level node
and n.get("file_type") != "rationale"
]
if not local_classes:
# Map each local symbol (class or function) to its node id, keyed by the
# bare symbol name. Function labels end in "()"; the file node ends in
# ".py"; rationale nodes never import (#563). First writer wins on a
# name collision (inherently ambiguous within a file).
name_to_nid: dict[str, str] = {}
for n in file_result.get("nodes", []):
if n.get("source_file") != str_path or n.get("file_type") == "rationale":
continue
label = n.get("label", "")
if not label or label.endswith(".py"):
continue
sym_name = label[:-2] if label.endswith("()") else label
if sym_name and sym_name not in name_to_nid:
name_to_nid[sym_name] = n["id"]
if not name_to_nid:
continue
# Parse imports from this file
@@ -1963,75 +1969,103 @@ def _resolve_cross_file_imports(
except Exception:
continue
def walk_imports(node) -> None:
if node.type == "import_from_statement":
# Find the module name - handles both absolute and relative imports.
# Relative: `from .models import X` → relative_import → dotted_name
# Absolute: `from models import X` → module_name field
# target_fq is the directory-qualified stem used as the key in
# stem_to_entities. Relative imports are resolved exactly via the
# importing file's directory; absolute imports fall back to the
# bare-stem secondary index (first-writer-wins when names collide).
target_fq: str | None = None
for child in node.children:
if child.type == "relative_import":
for sub in child.children:
if sub.type == "dotted_name":
raw = source[sub.start_byte:sub.end_byte].decode("utf-8", errors="replace")
bare = raw.split(".")[-1]
# Resolve relative import to exact qualified stem.
candidate = path.parent / f"{bare}.py"
target_fq = _file_stem(candidate)
break
break
if child.type == "dotted_name" and target_fq is None:
raw = source[child.start_byte:child.end_byte].decode("utf-8", errors="replace")
bare = raw.split(".")[-1]
target_fq = bare_to_qualified.get(bare)
# local_name -> target node id (local_name honours `import X as Y`, so a
# reference to the alias in the body still attributes correctly).
import_targets: dict[str, str] = {}
# referenced name -> {source symbol nid: first reference line}
ref_sources: dict[str, dict[str, int]] = {}
if not target_fq or target_fq not in stem_to_entities:
return
def _text(n) -> str:
return source[n.start_byte:n.end_byte].decode("utf-8", errors="replace")
# Collect imported names: dotted_name children of import_from_statement
# that come AFTER the 'import' keyword token.
imported_names: list[str] = []
past_import_kw = False
for child in node.children:
if child.type == "import":
past_import_kw = True
continue
if not past_import_kw:
continue
if child.type == "dotted_name":
imported_names.append(
source[child.start_byte:child.end_byte].decode("utf-8", errors="replace")
)
elif child.type == "aliased_import":
# `import X as Y` - take the original name
name_node = child.child_by_field_name("name")
if name_node:
imported_names.append(
source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
)
line = node.start_point[0] + 1
for name in imported_names:
tgt_nid = stem_to_entities[target_fq].get(name)
if tgt_nid:
for src_class_nid in local_classes:
new_edges.append({
"source": src_class_nid,
"target": tgt_nid,
"relation": "uses",
"confidence": "INFERRED",
"source_file": str_path,
"source_location": f"L{line}",
"weight": 0.8,
})
def resolve_import(node) -> None:
# Find the module name - handles both absolute and relative imports.
# Relative: `from .models import X` → relative_import → dotted_name
# Absolute: `from models import X` → module_name field
# target_fq is the directory-qualified stem used as the key in
# stem_to_entities. Relative imports are resolved exactly via the
# importing file's directory; absolute imports fall back to the
# bare-stem secondary index (first-writer-wins when names collide).
target_fq: str | None = None
for child in node.children:
walk_imports(child)
if child.type == "relative_import":
for sub in child.children:
if sub.type == "dotted_name":
bare = _text(sub).split(".")[-1]
candidate = path.parent / f"{bare}.py"
target_fq = _file_stem(candidate)
break
break
if child.type == "dotted_name" and target_fq is None:
bare = _text(child).split(".")[-1]
target_fq = bare_to_qualified.get(bare)
walk_imports(tree.root_node)
if not target_fq or target_fq not in stem_to_entities:
return
# Imported names come AFTER the 'import' keyword token. For
# `import X as Y` the target is found via X but the body uses Y.
past_import_kw = False
for child in node.children:
if child.type == "import":
past_import_kw = True
continue
if not past_import_kw:
continue
imported_name: str | None = None
local_name: str | None = None
if child.type == "dotted_name":
imported_name = local_name = _text(child)
elif child.type == "aliased_import":
name_node = child.child_by_field_name("name")
alias_node = child.child_by_field_name("alias")
if name_node is not None:
imported_name = _text(name_node)
local_name = _text(alias_node) if alias_node is not None else imported_name
if not imported_name or not local_name:
continue
tgt_nid = stem_to_entities[target_fq].get(imported_name)
if tgt_nid:
import_targets[local_name] = tgt_nid
def visit(node, current_nid: str | None) -> None:
# Identifiers inside an import statement are the import itself, not a
# real use — resolve the import here and don't descend into it.
if node.type == "import_from_statement":
resolve_import(node)
return
# Attribute references to the top-level symbol that contains them: a
# class is a unit (a reference inside one of its methods counts for
# the class, matching the documented DigestAuth->Response edge), and
# a module-level function is its own source. Only set at module scope
# (current_nid is None) so nested defs never override the container.
if current_nid is None and node.type in ("class_definition", "function_definition"):
name_node = node.child_by_field_name("name")
if name_node is not None:
mapped = name_to_nid.get(_text(name_node))
if mapped is not None:
current_nid = mapped
if node.type == "identifier" and current_nid is not None:
slot = ref_sources.setdefault(_text(node), {})
slot.setdefault(current_nid, node.start_point[0] + 1)
for child in node.children:
visit(child, current_nid)
visit(tree.root_node, None)
for name, tgt_nid in import_targets.items():
for src_nid, line in ref_sources.get(name, {}).items():
if src_nid == tgt_nid:
continue
new_edges.append({
"source": src_nid,
"target": tgt_nid,
"relation": "uses",
"confidence": "INFERRED",
"source_file": str_path,
"source_location": f"L{line}",
"weight": 0.8,
})
return new_edges
+60
View File
@@ -3434,3 +3434,63 @@ def test_extract_emits_posix_source_file_for_relative_inputs(tmp_path):
assert {sf for _, sf in carriers} == {
"src/lib/content.ts", "src/pages/index.astro",
}
def _inferred_uses(result):
"""(source, target) pairs of every INFERRED cross-file `uses` edge."""
return {
(e["source"], e["target"])
for e in result["edges"]
if e.get("relation") == "uses" and e.get("confidence") == "INFERRED"
}
def test_inferred_uses_edge_attributes_to_the_referencing_symbol(tmp_path):
"""A cross-file INFERRED `uses` edge binds to the symbol that actually
references the import — a function is a valid source and a co-located class
that never touches the import gets no edge (#2652)."""
(tmp_path / "helpers.py").write_text("class Helper:\n pass\n", encoding="utf-8")
(tmp_path / "api.py").write_text(
"from helpers import Helper\n\n\n"
"class Request:\n x: int = 0\n\n\n"
"def handler(req):\n return Helper()\n",
encoding="utf-8",
)
result = extract([tmp_path / "api.py", tmp_path / "helpers.py"], cache_root=tmp_path)
uses = _inferred_uses(result)
# handler() references Helper -> it is the source.
assert ("api_handler", "helpers_helper") in uses
# Request never references Helper -> no false edge from the co-located class.
assert ("api_request", "helpers_helper") not in uses
def test_inferred_uses_edge_kept_when_the_class_body_references_the_import(tmp_path):
"""Positive control: a class that genuinely uses the imported symbol still
gets its class-level INFERRED `uses` edge (the DigestAuth->Response case)."""
(tmp_path / "models.py").write_text("class Response:\n pass\n", encoding="utf-8")
(tmp_path / "auth.py").write_text(
"from models import Response\n\n\n"
"class DigestAuth:\n def build(self):\n return Response()\n",
encoding="utf-8",
)
result = extract([tmp_path / "auth.py", tmp_path / "models.py"], cache_root=tmp_path)
assert ("auth_digestauth", "models_response") in _inferred_uses(result)
def test_inferred_uses_edge_follows_an_import_alias(tmp_path):
"""`from helpers import Helper as H` attributes via the local alias `H`, so a
body that only ever names `H` still resolves to the imported target (#2652)."""
(tmp_path / "helpers.py").write_text("class Helper:\n pass\n", encoding="utf-8")
(tmp_path / "api.py").write_text(
"from helpers import Helper as H\n\n\n"
"def handler(req):\n return H()\n",
encoding="utf-8",
)
result = extract([tmp_path / "api.py", tmp_path / "helpers.py"], cache_root=tmp_path)
assert ("api_handler", "helpers_helper") in _inferred_uses(result)