Resolve cross-file Swift class relationships from member calls and constructors (#1356)

Capture property/field initializer constructor calls, build a per-file Swift
type table from property/parameter declarations, and add a member-call
resolution pass that types the receiver and emits an edge only when the type
name resolves to exactly one definition. Additive and INFERRED-only; the
is_member_call drop and the #543/#1219 god-node guards stay intact.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Safi
2026-06-17 12:08:59 +01:00
co-authored by Claude Opus 4.8
parent d8fa70e17b
commit 2e01f37c81
2 changed files with 359 additions and 1 deletions
+215 -1
View File
@@ -1090,6 +1090,57 @@ def _swift_property_type_node(property_node):
return None
def _swift_property_name(property_node, source: bytes) -> str | None:
"""Return the bound name of a Swift property (``let x``/``var x = ...``)."""
for c in property_node.children:
if c.type == "pattern":
for sc in c.children:
if sc.type == "simple_identifier":
return _read_text(sc, source)
if c.type == "simple_identifier":
return _read_text(c, source)
return None
def _swift_constructor_type(call_node, source: bytes) -> str | None:
"""If a Swift call expression is a constructor (``Foo()``), return the type name.
Only upper-cased callees are treated as types so a free-function call like
``configure()`` in an initializer is not mistaken for a constructor.
"""
first = call_node.children[0] if call_node.children else None
if first is not None and first.type == "simple_identifier":
text = _read_text(first, source)
if text and text[:1].isupper():
return text
return None
def _swift_receiver_name(recv_node, source: bytes) -> str | None:
"""Return the depth-1 receiver name of a Swift member call (``recv.method()``).
``vm.update()`` -> ``vm``; ``Type.staticMethod()`` -> ``Type``;
``Singleton.shared.method()`` -> ``Singleton`` (head of the chain);
``self.svc.fetch()`` -> ``svc`` (the property the call is reached through).
Returns None for anything deeper, so resolution stays depth-1.
"""
if recv_node is None:
return None
if recv_node.type == "simple_identifier":
return _read_text(recv_node, source)
if recv_node.type == "navigation_expression":
head = recv_node.children[0] if recv_node.children else None
if head is not None and head.type == "simple_identifier":
return _read_text(head, source)
if head is not None and head.type == "self_expression":
for child in recv_node.children:
if child.type == "navigation_suffix":
for sc in child.children:
if sc.type == "simple_identifier":
return _read_text(sc, source)
return None
# ── C / C++ type-ref helpers ─────────────────────────────────────────────────
_C_PRIMITIVE_TYPE_NODES = frozenset({
@@ -2363,6 +2414,14 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
# extensions don't (file stem is part of the id), so they're collected here
# for a corpus-level merge after every file has been parsed.
swift_extensions: list[dict] = []
# #1356: call expressions in property/field initializers (e.g.
# `let vm = VM()`) live outside function bodies, so the call-walk never
# reaches them. Collect (owner_nid, call_node) here and walk them too.
initializer_nodes: list[tuple[str, object]] = []
# #1356: per-file map of local name -> declared type (properties + params),
# threaded out as `swift_type_table` so member calls (`vm.update()`) can be
# resolved to the receiver's real definition in _resolve_swift_member_calls.
type_table: dict[str, str] = {}
csharp_interface_names: set[str] = set()
if config.ts_module == "tree_sitter_c_sharp":
@@ -2945,9 +3004,10 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
if (config.ts_module == "tree_sitter_swift"
and t == "property_declaration"
and parent_class_nid):
line = node.start_point[0] + 1
prop_type: str | None = None
type_anno = _swift_property_type_node(node)
if type_anno is not None:
line = node.start_point[0] + 1
refs: list[tuple[str, str]] = []
_swift_collect_type_refs(type_anno, source, False, refs)
for ref_name, role in refs:
@@ -2955,6 +3015,22 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
target_nid = ensure_named_node(ref_name, line)
if target_nid != parent_class_nid:
add_edge(parent_class_nid, target_nid, "references", line, context=ctx)
if prop_type is None and role == "type":
prop_type = ref_name
# #1356 Stage 1: walk the initializer so a constructor call
# (`let vm = VM()`) produces a calls edge. #1356 Stage 2a: when the
# property has no type annotation, infer its type from the
# constructor so `vm.update()` later resolves to VM.
for child in node.children:
if child.type in config.call_types:
initializer_nodes.append((parent_class_nid, child))
if prop_type is None:
ctor = _swift_constructor_type(child, source)
if ctor is not None:
prop_type = ctor
prop_name = _swift_property_name(node, source)
if prop_name and prop_type:
type_table[prop_name] = prop_type
return
if (config.ts_module == "tree_sitter_scala"
@@ -3195,11 +3271,22 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
type_node = p.child_by_field_name("type")
refs: list[tuple[str, str]] = []
_swift_collect_type_refs(type_node, source, False, refs)
param_type: str | None = None
for ref_name, role in refs:
ctx = "generic_arg" if role == "generic_arg" else "parameter_type"
target_nid = ensure_named_node(ref_name, line)
if target_nid != func_nid:
add_edge(func_nid, target_nid, "references", line, context=ctx)
if param_type is None and role == "type":
param_type = ref_name
# #1356 Stage 2a: record param name -> type (flat per-file
# table; later params with the same name win, which is fine
# for the depth-1 member-call resolution we do).
if param_type:
name_node = p.child_by_field_name("name")
pname = _read_text(name_node, source) if name_node else None
if pname:
type_table[pname] = param_type
return_node = node.child_by_field_name("return_type")
if return_node is not None:
refs = []
@@ -3397,6 +3484,7 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
callee_name: str | None = None
is_member_call: bool = False
swift_receiver: str | None = None
# Special handling per language
if config.ts_module == "tree_sitter_swift":
@@ -3412,6 +3500,10 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
for sc in child.children:
if sc.type == "simple_identifier":
callee_name = _read_text(sc, source)
# #1356: capture the receiver so the cross-file pass can
# resolve it through the file's type table.
recv_node = first.children[0] if first.children else None
swift_receiver = _swift_receiver_name(recv_node, source)
elif config.ts_module == "tree_sitter_kotlin":
# Kotlin: first child may be simple_identifier/identifier or
# navigation_expression. PyPI's `tree_sitter_kotlin` produces
@@ -3528,6 +3620,7 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
"is_member_call": is_member_call,
"source_file": str_path,
"source_location": f"L{node.start_point[0] + 1}",
"receiver": swift_receiver,
})
# Helper function calls: config('foo.bar') → uses_config edge to "foo"
@@ -3660,6 +3753,12 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
for caller_nid, body_node in function_bodies:
walk_calls(body_node, caller_nid)
# #1356: walk property/field initializers (collected above). walk_calls
# self-guards against re-entering function bodies and dedups via
# seen_call_pairs, so a closure inside an initializer is not double-walked.
for owner_nid, init_node in initializer_nodes:
walk_calls(init_node, owner_nid)
# ── Event listener pass ───────────────────────────────────────────────────
seen_listen_pairs: set[tuple[str, str]] = set()
for event_name, listener_name, line in pending_listen_edges:
@@ -3693,6 +3792,8 @@ def _extract_generic(path: Path, config: LanguageConfig) -> dict:
result = {"nodes": nodes, "edges": clean_edges, "raw_calls": raw_calls}
if swift_extensions:
result["swift_extensions"] = swift_extensions
if type_table:
result["swift_type_table"] = {"path": str_path, "table": type_table}
return result
@@ -9040,6 +9141,108 @@ def _resolve_java_type_references(
]
def _resolve_swift_member_calls(
per_file: list[dict],
all_nodes: list[dict],
all_edges: list[dict],
) -> None:
"""Resolve cross-file Swift member calls (``recv.method()``) to the real
definition of the receiver's type (#1356).
The shared cross-file call pass drops every ``is_member_call`` because a bare
method name (``update``) collides across the corpus and inflates god-nodes
(#543/#1219). Swift extractors record the receiver of each member call and a
per-file ``name -> type`` table (``swift_type_table``); this pass uses them to
type the receiver, then emits an edge ONLY when that type name resolves to
exactly one definition. Everything it adds is INFERRED (type inference, not an
explicit import), and the line-12503 drop stays intact: this is purely
additive and fires only on receiver-typed Swift calls.
Must run after id-disambiguation so node ids and caller_nids are final.
"""
type_table_by_file: dict[str, dict[str, str]] = {}
for result in per_file:
tt = result.get("swift_type_table")
if tt and tt.get("path"):
type_table_by_file[tt["path"]] = tt.get("table", {})
if not type_table_by_file:
return
def _key(label: str) -> str:
return re.sub(r"[^a-zA-Z0-9]+", "", str(label)).lower()
# A genuine Swift type is the target of a `contains` edge from its file node.
# Bare type references create a same-label shadow node (via ensure_named_node)
# that carries a source_file but is NOT contained; excluding non-contained
# nodes keeps that shadow from making a real type name look ambiguous.
contained = {e.get("target") for e in all_edges if e.get("relation") == "contains"}
# Type name -> definition node ids (real, source-backed, type-like defs only).
# len != 1 is the god-node guard: an ambiguous type name bails.
type_def_nids: dict[str, list[str]] = {}
node_by_id: dict[str, dict] = {}
for n in all_nodes:
node_by_id[n.get("id")] = n
if n.get("source_file") and n.get("id") in contained and _is_type_like_definition(n):
type_def_nids.setdefault(_key(n.get("label", "")), []).append(n["id"])
# (type_node_id, method_key) -> method_node_id, from `method` edges.
method_index: dict[tuple[str, str], str] = {}
for e in all_edges:
if e.get("relation") != "method":
continue
src, tgt = e.get("source"), e.get("target")
tnode = node_by_id.get(tgt)
if tnode is not None:
method_index[(src, _key(tnode.get("label", "")))] = tgt
all_raw_calls: list[dict] = []
for result in per_file:
all_raw_calls.extend(result.get("raw_calls", []))
existing_pairs = {(e.get("source"), e.get("target")) for e in all_edges}
for rc in all_raw_calls:
if not rc.get("is_member_call"):
continue
receiver = rc.get("receiver")
callee = rc.get("callee")
if not receiver or not callee:
continue
# Determine the receiver's type. An upper-cased receiver is itself a type
# (Type.staticMethod(), Singleton.shared.x()); otherwise look it up in the
# declaring file's local type table.
if receiver[:1].isupper():
type_name = receiver
else:
type_name = type_table_by_file.get(rc.get("source_file", ""), {}).get(receiver)
if not type_name:
continue
type_defs = type_def_nids.get(_key(type_name), [])
if len(type_defs) != 1: # ambiguous or absent -> bail (god-node guard)
continue
type_nid = type_defs[0]
caller = rc.get("caller_nid")
if not caller:
continue
method_nid = method_index.get((type_nid, _key(callee)))
target = method_nid or type_nid
relation = "calls" if method_nid else "references"
if target == caller or (caller, target) in existing_pairs:
continue
existing_pairs.add((caller, target))
all_edges.append({
"source": caller,
"target": target,
"relation": relation,
"context": "call",
"confidence": "INFERRED",
"confidence_score": 0.8,
"source_file": rc.get("source_file", ""),
"source_location": rc.get("source_location"),
"weight": 1.0,
})
def extract_objc(path: Path) -> dict:
"""Extract interfaces, implementations, protocols, methods, and imports from .m/.mm/.h files."""
try:
@@ -12568,6 +12771,17 @@ def extract(
"weight": 1.0,
})
# Cross-file Swift member-call resolution (#1356). Runs after the shared call
# pass so node ids/caller_nids are final; additive (only receiver-typed calls
# the shared pass skipped), with a single-definition god-node guard.
swift_paths = [p for p in paths if p.suffix == ".swift"]
if swift_paths:
try:
_resolve_swift_member_calls(per_file, all_nodes, all_edges)
except Exception as exc:
import logging
logging.getLogger(__name__).warning("Swift member-call resolution failed, skipping: %s", exc)
# Relativize source_file fields so paths are portable across machines (#555)
for item in all_nodes + all_edges:
sf = item.get("source_file")
+144
View File
@@ -0,0 +1,144 @@
from __future__ import annotations
from pathlib import Path
from graphify.build import build_from_json
from graphify.extract import extract
def _write(path: Path, text: str) -> Path:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(text, encoding="utf-8")
return path
def _label(result: dict, nid: str) -> str:
for n in result["nodes"]:
if n["id"] == nid:
return n.get("label", "")
return f"<{nid}>"
def _edge_labels(result: dict, relations=("calls", "references")) -> set[tuple[str, str, str]]:
"""Return {(source_label, relation, target_label)} for the given relations."""
out: set[tuple[str, str, str]] = set()
for e in result["edges"]:
if e.get("relation") in relations:
out.add((_label(result, e["source"]), e["relation"], _label(result, e["target"])))
return out
def _issue_fixture(base: Path) -> list[Path]:
"""The three cross-file patterns from #1356, plus a constructor-in-initializer."""
f1 = _write(base / "Models/SessionViewModel.swift",
"class SessionViewModel {\n func update() {}\n}\n")
f2 = _write(base / "Services/NetworkService.swift",
"class NetworkService {\n func fetch() {}\n}\n")
f3 = _write(base / "Core/SessionType.swift",
"enum SessionType {\n static func staticMethod() {}\n}\n")
f4 = _write(base / "Core/Singleton.swift",
"class Singleton {\n static let shared = Singleton()\n func method() {}\n}\n")
f5 = _write(base / "Views/HomeView.swift", (
"class HomeView {\n"
" let vm = SessionViewModel()\n"
" var svc: NetworkService\n\n"
" func go() {\n"
" vm.update()\n"
" SessionType.staticMethod()\n"
" Singleton.shared.method()\n"
" self.svc.fetch()\n"
" }\n"
"}\n"
))
return [f1, f2, f3, f4, f5]
def test_swift_cross_file_member_calls_resolve(tmp_path: Path):
# #1356: cross-file member calls (recv.method()), static/singleton calls, and
# a constructor-in-initializer must resolve to the receiver's real definition.
files = _issue_fixture(tmp_path / "src")
result = extract(files, cache_root=tmp_path / "cache")
edges = _edge_labels(result)
# Stage 1: constructor in a property initializer.
assert ("HomeView", "calls", "SessionViewModel") in edges
# Stage 2: receiver typed via the file's local type table.
assert (".go()", "calls", ".update()") in edges # vm.update()
assert (".go()", "calls", ".fetch()") in edges # self.svc.fetch()
# Stage 2: upper-cased receiver is itself a type.
assert (".go()", "calls", ".staticMethod()") in edges # SessionType.staticMethod()
assert (".go()", "calls", ".method()") in edges # Singleton.shared.method()
def test_swift_cross_file_member_calls_are_inferred_and_resolve_to_real_nodes(tmp_path: Path):
# The new edges must be INFERRED (type inference, not an explicit import) and
# land on real definition nodes so build_from_json keeps them.
files = _issue_fixture(tmp_path / "src")
result = extract(files, cache_root=tmp_path / "cache")
node_ids = {n["id"] for n in result["nodes"]}
src_by_id = {n["id"]: n.get("source_file") for n in result["nodes"]}
member_targets = {".update()", ".fetch()", ".staticMethod()", ".method()"}
seen_targets: set[str] = set()
for e in result["edges"]:
tgt_label = _label(result, e["target"])
if e.get("relation") == "calls" and tgt_label in member_targets:
assert e["confidence"] == "INFERRED"
assert e["confidence_score"] == 0.8
assert e["target"] in node_ids
assert src_by_id.get(e["target"]) # resolved to a real, source-backed def
seen_targets.add(tgt_label)
assert seen_targets == member_targets
# Edges survive graph construction (no dangling targets pruned).
g = build_from_json(result)
surviving = sum(
1 for _, _, d in g.edges(data=True)
if d.get("confidence") == "INFERRED" and d.get("relation") == "calls"
)
assert surviving >= 5
def test_swift_ambiguous_type_does_not_over_connect(tmp_path: Path):
# #543/#1219 guard: when the receiver's type name is defined in 2+ files the
# resolution must bail rather than fan a member call out to every candidate.
base = tmp_path / "src"
for sub in ("a", "b", "c"):
_write(base / sub / "Widget.swift", "class Widget {\n func update() {}\n}\n")
_write(base / "Caller.swift", (
"class Caller {\n"
" var w: Widget\n"
" func run() {\n"
" w.update()\n"
" unknown.update()\n"
" }\n"
"}\n"
))
files = sorted(base.rglob("*.swift"))
result = extract(files, cache_root=tmp_path / "cache")
inferred_calls = [
e for e in result["edges"]
if e.get("relation") == "calls" and e.get("confidence") == "INFERRED"
]
# Ambiguous `Widget` (3 defs) -> no member-call edge; unknown receiver -> none.
assert inferred_calls == []
def test_swift_unknown_receiver_emits_no_edge(tmp_path: Path):
# A lowercase receiver absent from the file's type table is never guessed.
base = tmp_path / "src"
_write(base / "Helper.swift", "class Helper {\n func help() {}\n}\n")
_write(base / "Caller.swift", (
"class Caller {\n"
" func run() {\n"
" mystery.help()\n"
" }\n"
"}\n"
))
files = sorted(base.rglob("*.swift"))
result = extract(files, cache_root=tmp_path / "cache")
edges = _edge_labels(result, relations=("calls",))
assert (".run()", "calls", ".help()") not in edges