diff --git a/graphify/extractors/resolution.py b/graphify/extractors/resolution.py index aa57e35..8af9aac 100644 --- a/graphify/extractors/resolution.py +++ b/graphify/extractors/resolution.py @@ -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 diff --git a/tests/test_extract.py b/tests/test_extract.py index cb9d715..e514205 100644 --- a/tests/test_extract.py +++ b/tests/test_extract.py @@ -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)