diff --git a/graphify/cache.py b/graphify/cache.py index f9031f7597..0e3254cbc1 100644 --- a/graphify/cache.py +++ b/graphify/cache.py @@ -36,7 +36,7 @@ _EXTRACTOR_VERSION = "unknown" # Bump when AST cache-key semantics change independently of the package version. -_AST_CACHE_SCHEMA = 5 # Python receiver-shadow facts in persisted raw calls. +_AST_CACHE_SCHEMA = 6 # Python typed-receiver facts (receiver_type) in persisted raw calls. # Version dirs already swept this process — cleanup runs once per (base, version). _cleaned_ast_dirs: set[str] = set() diff --git a/graphify/extract.py b/graphify/extract.py index db4fd1bad6..cf4e56a897 100644 --- a/graphify/extract.py +++ b/graphify/extract.py @@ -4199,21 +4199,39 @@ def _module_stem_key(nid: str) -> str: return _key(stem or n.get("label", "")) existing_pairs = {(e.get("source"), e.get("target")) for e in all_edges} + # A method hangs off its class (`method` edge), not its file (`contains`), so + # the typed-receiver arm finds a method caller's file through its class. + method_class: dict[str, str] = {} + for e in all_edges: + src, tgt = e.get("source"), e.get("target") + if e.get("relation") == "method" and isinstance(src, str) and isinstance(tgt, str): + method_class[tgt] = src + + def _file_node(nid: str) -> "str | None": + seen: set[str] = set() + while nid and nid not in seen: + seen.add(nid) + if nid in file_of_node: + return file_of_node[nid] + nid = method_class.get(nid, "") + return None - def _emit_call(caller: str, target_nid: "str | None", rc: dict) -> None: + def _emit_call(caller: str, target_nid: "str | None", rc: dict, confidence: str = "EXTRACTED") -> None: if not target_nid or target_nid == caller or (caller, target_nid) in existing_pairs: return existing_pairs.add((caller, target_nid)) # EXTRACTED: a qualified call (`ClassName.method()` or `module.func()`) is # an explicit, unambiguous static reference resolved to exactly one # definition (each arm applies a single-definition god-node guard). + # INFERRED: the typed-receiver arm, where the receiver's class came from + # an annotation or a constructor binding rather than the call itself. all_edges.append({ "source": caller, "target": target_nid, "relation": "calls", "context": "call", - "confidence": "EXTRACTED", - "confidence_score": 1.0, + "confidence": confidence, + "confidence_score": 1.0 if confidence == "EXTRACTED" else 0.85, # rubric value (#2813) "source_file": rc.get("source_file", ""), "source_location": rc.get("source_location"), "weight": 1.0, @@ -4241,6 +4259,25 @@ def _emit_call(caller: str, target_nid: "str | None", rc: dict) -> None: continue if targets: continue + receiver_type = rc.get("receiver_type") + if receiver_type and not receiver[:1].isupper(): + # Typed-receiver arm: `request.read()` where the extractor typed the + # local or parameter `request` from an annotation (`request: Request`) + # or a constructor binding (`client = Client(...)`, `with Client() as + # client`). Same origin gate as the TypeScript arm (#2553): the class + # must be the only class of that name AND be defined in the caller's + # file, imported by name into it, or inside a module it imports; + # otherwise emit nothing. The method must be the class's own. A typed + # local is never an imported module, so the module arm is skipped. + class_nids = class_def_nids.get(_key(receiver_type), []) + caller_file = _file_node(caller) + if len(class_nids) == 1 and caller_file is not None: + cls = class_nids[0] + imported = imported_by_filenode.get(caller_file, set()) + cls_file = _file_node(cls) + if cls_file == caller_file or cls in imported or (cls_file is not None and cls_file in imported): + _emit_call(caller, method_index.get((cls, _key(callee))), rc, confidence="INFERRED") + continue if receiver[:1].isupper(): # Class arm (#1446): a capitalized receiver is a class reference; an # instance (`self`, `obj`) never collides with a same-spelled class. diff --git a/graphify/extractors/engine.py b/graphify/extractors/engine.py index b07106585a..6a6a47d6fb 100644 --- a/graphify/extractors/engine.py +++ b/graphify/extractors/engine.py @@ -3691,6 +3691,172 @@ def _ruby_new_class_name(node, source: bytes) -> str | None: return None return _read_text(recv, source) +def _python_class_ref(node, source: bytes) -> str | None: + """The single class a Python annotation or constructor callee names, else None. + + ``Request``, ``httpx.Request``, ``"Request"``, ``Optional[Request]`` and + ``Request | None`` all name ``Request``. A container (``list[Request]``), a + union of two classes, or a lowercase callee (``make_client``, + ``Client.from_url``) names no single class, so it stays untyped. + """ + if node is None: + return None + kind = node.type + if kind == "type": + named = [c for c in node.children if c.is_named] + return _python_class_ref(named[0], source) if len(named) == 1 else None + if kind == "identifier": + name = _read_text(node, source) + return name if name[:1].isupper() else None + if kind == "attribute": + return _python_class_ref(node.child_by_field_name("attribute"), source) + if kind == "string": + name = _read_text(node, source).strip("\"'").split(".")[-1].strip() + return name if name.isidentifier() and name[:1].isupper() else None + if kind == "subscript": # typing.Optional[X] + value = node.child_by_field_name("value") + subs = node.children_by_field_name("subscript") + if value is not None and _read_text(value, source).split(".")[-1] == "Optional" and len(subs) == 1: + return _python_class_ref(subs[0], source) + return None + if kind == "generic_type": # Optional[X] + named = [c for c in node.children if c.is_named] + if ( + len(named) == 2 + and named[0].type == "identifier" + and _read_text(named[0], source) == "Optional" + and named[1].type == "type_parameter" + ): + inner = [c for c in named[1].children if c.is_named] + return _python_class_ref(inner[0], source) if len(inner) == 1 else None + return None + if kind == "binary_operator": + left, right = node.child_by_field_name("left"), node.child_by_field_name("right") + if right is not None and right.type == "none": + return _python_class_ref(left, source) + if left is not None and left.type == "none": + return _python_class_ref(right, source) + return None + +def _python_local_class_bindings(body_node, source: bytes) -> dict[str, str | None]: + """Map ``name -> ClassName`` for the parameters and locals of one Python function. + + Typed from a class annotation on a parameter or local (``request: Request``, + ``x: Client = ...``), a constructor call (``client = Client(...)``) or a + context manager (``with Client(...) as client``). Same 100%-confidence + contract as the Ruby table: the table is flow-insensitive, so a name bound + to anything that is not one class -- an untyped parameter, a factory call, a + loop target, a second class, tuple unpacking, a lambda parameter, a ``match`` + capture, an import -- maps to ``None`` for the whole function and is never + resolved. Nested functions are their own caller and are not entered. Calls + inside a lambda or a nested class body are attributed to this function, so + every name those bind is poisoned rather than typed. + """ + bindings: dict[str, str | None] = {} + poison_only = False + + def bind(name: str, cls: str | None) -> None: + if poison_only or cls is None or bindings.get(name, cls) != cls: + bindings[name] = None + elif name not in bindings: + bindings[name] = cls + + def poison(target) -> None: + names: set[str] = set() + _python_collect_assignment_targets(target, source, names) + for name in names: + bindings[name] = None + + def poison_identifiers(node) -> None: + stack = [node] + while stack: + cur = stack.pop() + if cur.type == "identifier": + bindings[_read_text(cur, source)] = None + stack.extend(cur.children) + + func = body_node.parent + params = func.child_by_field_name("parameters") if func is not None and func.type == "function_definition" else None + for name in _python_param_names(params, source): + bindings[name] = None # an untyped parameter: unknown, and a later assignment must not type it + if params is not None: + for child in params.children: + if child.type in ("typed_parameter", "typed_default_parameter"): + name_n = child.child_by_field_name("name") or next( + (c for c in child.children if c.type == "identifier"), None + ) + if name_n is not None: + name = _read_text(name_n, source) + bindings.pop(name, None) + bind(name, _python_class_ref(child.child_by_field_name("type"), source)) + + def visit(n) -> None: + nonlocal poison_only + for child in n.children: + if child.type == "function_definition": + # Its own caller, but `nonlocal c` inside it can rebind this function's `c`. + stack = [child] + while stack: + cur = stack.pop() + if cur.type == "nonlocal_statement": + poison_identifiers(cur) + stack.extend(cur.children) + continue + if child.type in ("lambda", "class_definition"): + # Their own scope, but their calls are still this function's: a name + # they bind (`lambda c: c.send()`, `class K: c = Cache()`) must not + # keep the outer type, and must not gain one either. + params = child.child_by_field_name("parameters") + if child.type == "lambda" and params is not None: + poison_identifiers(params) + outer, poison_only = poison_only, True + visit(child) + poison_only = outer + continue + if child.type in ("case_pattern", "import_statement", "import_from_statement"): + poison_identifiers(child) # `case c:` / `import c` rebind the name + continue + if child.type == "assignment": + left = child.child_by_field_name("left") + if left is not None and left.type == "identifier": + name = _read_text(left, source) + annotation = child.child_by_field_name("type") + right = child.child_by_field_name("right") + if annotation is not None: + bind(name, _python_class_ref(annotation, source)) + elif right is not None and right.type == "call": + bind(name, _python_class_ref(right.child_by_field_name("function"), source)) + else: + bindings[name] = None + else: + poison(left) + elif child.type == "augmented_assignment": + poison(child.child_by_field_name("left")) + elif child.type == "as_pattern": + alias = child.child_by_field_name("alias") + target = next((c for c in alias.children if c.type == "identifier"), None) if alias is not None else None + value = next((c for c in child.children if c.is_named), None) + if target is not None: + name = _read_text(target, source) + if child.parent is not None and child.parent.type == "with_item" and value is not None and value.type == "call": + bind(name, _python_class_ref(value.child_by_field_name("function"), source)) + else: + bindings[name] = None # `except E as e`, other patterns: not a constructor result + elif alias is not None: + poison_identifiers(alias) # `with A() as (x, y)` + elif child.type in ("for_statement", "for_in_clause"): + poison(child.child_by_field_name("left")) + elif child.type == "named_expression": + poison(child.child_by_field_name("name")) + elif child.type in ("global_statement", "nonlocal_statement"): + for c in child.children: + if c.type == "identifier": + bindings[_read_text(c, source)] = None + visit(child) + + visit(body_node) + return bindings + def _ruby_local_class_bindings(body_node, source: bytes) -> dict[str, str | None]: """Map ``local_var -> ClassName`` for ``var = ClassName.new`` within one Ruby method body, not descending into nested method definitions. @@ -6638,6 +6804,9 @@ def scala_base_name(type_node) -> str | None: # populated before walk_calls runs. Lets member-call raw_calls carry a # receiver_type so the cross-file pass resolves `var.method` by type (#ruby). ruby_var_types: dict[str, dict[str, str | None]] = {} + # Python: per-function `name -> ClassName` for typed parameters and locals, + # stamped on member calls so the corpus resolver can type `obj.method()`. + python_var_types: dict[str, dict[str, str | None]] = {} # Ruby: per-method set of bound local/parameter names, so the call-walk can # tell a paren-less self-send (`build`) from a plain variable read (`x`). ruby_local_names: dict[str, frozenset[str]] = {} @@ -7589,6 +7758,11 @@ def walk_calls( rc_entry["receiver_type"] = ruby_var_types.get( caller_nid, {} ).get(member_receiver) + # Python: the same, from typed parameters and locals. + if member_receiver and config.ts_module == "tree_sitter_python": + _py_type = python_var_types.get(caller_nid, {}).get(member_receiver) + if _py_type: + rc_entry["receiver_type"] = _py_type # Tag the C++ raw_call's language so the cross-file C++ resolver # claims it unambiguously: a `.h` file routes to extract_cpp or # extract_objc by content, and both resolvers see `.h` in their @@ -7894,6 +8068,10 @@ def walk_calls( ruby_var_types[caller_nid] = _ruby_local_class_bindings(body_node, source) ruby_local_names[caller_nid] = _ruby_local_names(body_node, source) + if config.ts_module == "tree_sitter_python": + for caller_nid, body_node in function_bodies: + python_var_types[caller_nid] = _python_local_class_bindings(body_node, source) + # C++: build the per-file `var -> ClassName` table from local declarations in # every function body so the cross-file member-call pass can type a receiver # (#1547). File-scoped (not per-body): a later body's `Foo f;` doesn't clobber diff --git a/tests/test_python_typed_receiver_calls.py b/tests/test_python_typed_receiver_calls.py new file mode 100644 index 0000000000..aac3748a59 --- /dev/null +++ b/tests/test_python_typed_receiver_calls.py @@ -0,0 +1,204 @@ +"""Python member calls on a receiver whose class is known from the function itself. + +Before this, `obj.method()` produced a `calls` edge only for `self`/`cls` receivers and +module aliases, so `affected Client.send` missed every caller that wrote + + A. `def f(client: Client): client.send()` (annotated parameter, incl. Optional/str/`| None`) + B. `client = Client(); client.send()` (local bound from a constructor call) + C. `with Client() as client: client.send()` (context manager on a constructor call) + +The receiver's class is used only when the binding is unambiguous inside the function +(one class, never rebound to something else) and the class is unique by name and visible +from the caller's file (same file or imported), mirroring the TS receiver gate (#2553). +Everything else emits no edge rather than a guess. +""" +from __future__ import annotations + +from graphify.extract import extract + +_CLIENT = "class Client:\n def send(self):\n return 1\n" +_CACHE = "class Cache:\n def send(self):\n return 2\n" + + +def _calls(tmp_path, files: dict[str, str]): + # A real package, so `from .client import Client` binds to the class node. + files = {"pkg/__init__.py": "", **{f"pkg/{n}": b for n, b in files.items()}} + for name, body in files.items(): + p = tmp_path / name + p.parent.mkdir(parents=True, exist_ok=True) + p.write_text(body, encoding="utf-8") + r = extract([tmp_path / n for n in files], + cache_root=tmp_path / "graphify-out", parallel=False) + lbl = {n["id"]: n["label"] for n in r["nodes"]} + edges = [e for e in r["edges"] if e["relation"] == "calls"] + return {(lbl.get(e["source"]), lbl.get(e["target"])) for e in edges}, edges, r + + +def _has(calls, caller, callee=".send()"): + return any(s == f"{caller}()" and t == callee for s, t in calls) + + +def test_annotated_parameter_receiver(tmp_path): + calls, edges, _ = _calls(tmp_path, { + "client.py": _CLIENT, + "use.py": "from .client import Client\n\ndef run(c: Client):\n return c.send()\n", + }) + assert _has(calls, "run") + edge = next(e for e in edges if e["source"].endswith("_run")) + assert edge["confidence"] == "INFERRED" + assert edge["confidence_score"] == 0.85 + + +def test_optional_string_and_union_annotations(tmp_path): + calls, _, _ = _calls(tmp_path, { + "client.py": _CLIENT, + "use.py": ( + "from typing import Optional\nfrom .client import Client\n\n" + "def opt(c: Optional[Client]):\n return c.send()\n\n" + "def fwd(c: 'Client'):\n return c.send()\n\n" + "def union(c: Client | None = None):\n return c.send()\n" + ), + }) + assert _has(calls, "opt") and _has(calls, "fwd") and _has(calls, "union") + + +def test_local_constructor_and_with_binding(tmp_path): + calls, _, _ = _calls(tmp_path, { + "client.py": _CLIENT, + "use.py": ( + "from .client import Client\n\n" + "def local():\n c = Client()\n return c.send()\n\n" + "def ctx():\n with Client() as c:\n return c.send()\n" + ), + }) + assert _has(calls, "local") and _has(calls, "ctx") + + +def test_module_qualified_typing_wrapper(tmp_path): + calls, _, _ = _calls(tmp_path, { + "client.py": _CLIENT, + "use.py": ("import typing\nfrom .client import Client\n\n" + "def run(c: typing.Optional[Client]):\n return c.send()\n"), + }) + assert _has(calls, "run") + + +def test_same_method_name_resolves_to_annotated_class(tmp_path): + _, edges, r = _calls(tmp_path, { + "client.py": _CLIENT, + "cache.py": _CACHE, + "use.py": "from .client import Client\n\ndef run(c: Client):\n return c.send()\n", + }) + files = {n["id"]: str(n.get("source_file", "")) for n in r["nodes"]} + targets = [files[e["target"]] for e in edges if e["source"].endswith("_run")] + assert targets and all(t.endswith("client.py") for t in targets) + + +def test_unknown_or_ambiguous_receivers_emit_no_edge(tmp_path): + calls, _, _ = _calls(tmp_path, { + "client.py": _CLIENT, + "cache.py": _CACHE, + "use.py": ( + "from .client import Client\nfrom .cache import Cache\n\n" + "def untyped(c):\n return c.send()\n\n" + "def rebound(c: Client):\n c = Cache()\n return c.send()\n\n" + "def from_func():\n c = make()\n return c.send()\n\n" + "def container(cs: list[Client]):\n return cs.send()\n\n" + "def loop(cs):\n for c in cs:\n c.send()\n\n" + "def caught():\n try:\n pass\n except Exception as c:\n c.send()\n" + ), + }) + for caller in ("untyped", "rebound", "from_func", "container", "loop", "caught"): + assert not _has(calls, caller), caller + + +def test_rebinding_in_an_inner_scope_emits_no_edge(tmp_path): + # Calls inside a lambda or a nested class body belong to the enclosing function, + # so a name those scopes (or match/import/nonlocal) bind must lose its type. + calls, _, _ = _calls(tmp_path, { + "client.py": _CLIENT, + "cache.py": _CACHE, + "use.py": ( + "from .client import Client\nfrom .cache import Cache\n\n" + "def lam(c: Client):\n return lambda c: c.send()\n\n" + "def klass(c: Client):\n class K:\n c = Cache()\n c.send()\n return K\n\n" + "def matched(c: Client, v):\n match v:\n case c:\n return c.send()\n\n" + "def imported(c: Client):\n import c\n return c.send()\n\n" + "def nonloc(c: Client):\n def inner():\n nonlocal c\n c = Cache()\n" + " inner()\n return c.send()\n\n" + "def unpacked(c: Client):\n with Cache() as (c, d):\n return c.send()\n\n" + "def star(*c: Client):\n return c.send()\n\n" + "def other_lambda(c: Client):\n f = lambda x: x\n return c.send()\n" + ), + }) + for caller in ("lam", "klass", "matched", "imported", "nonloc", "unpacked", "star"): + assert not _has(calls, caller), caller + assert _has(calls, "other_lambda") + + +def test_class_not_visible_from_caller_emits_no_edge(tmp_path): + calls, _, _ = _calls(tmp_path, { + "client.py": _CLIENT, + "use.py": "def run(c: 'Client'):\n return c.send()\n", + }) + assert not _has(calls, "run") + + +def test_duplicate_class_name_emits_no_edge(tmp_path): + calls, _, _ = _calls(tmp_path, { + "a/__init__.py": "", + "a/client.py": _CLIENT, + "b/__init__.py": "", + "b/client.py": _CLIENT, + "use.py": "from .a.client import Client\n\ndef run(c: Client):\n return c.send()\n", + }) + assert not _has(calls, "run") + + +def test_schema_bump_retires_raw_calls_cached_without_receiver_type(tmp_path, monkeypatch): + """A cache written before receiver_type existed must re-extract, not replay a + raw call the typed-receiver arm cannot resolve (same shape as the Rust test).""" + import graphify.cache as cache_mod + from graphify.cache import save_cached + from graphify.extract import extract_python + + files = { + "pkg/__init__.py": "", + "pkg/client.py": _CLIENT, + "pkg/use.py": "from .client import Client\n\ndef run(c: Client):\n return c.send()\n", + } + paths = [] + for name, body in files.items(): + p = tmp_path / name + p.parent.mkdir(parents=True, exist_ok=True) + p.write_text(body, encoding="utf-8") + paths.append(p) + + current_schema = cache_mod._AST_CACHE_SCHEMA + monkeypatch.setattr(cache_mod, "_EXTRACTOR_VERSION", "same-version") + monkeypatch.setattr(cache_mod, "_AST_CACHE_SCHEMA", 5) # last schema without receiver_type + monkeypatch.setattr(cache_mod, "_cleaned_ast_dirs", set()) + for p in paths: + stale = extract_python(p) + for raw_call in stale.get("raw_calls", []): + raw_call.pop("receiver_type", None) # what the previous schema persisted + save_cached(p, stale, root=tmp_path, cache_root=tmp_path, kind="ast") + + monkeypatch.setattr(cache_mod, "_AST_CACHE_SCHEMA", current_schema) + monkeypatch.setattr(cache_mod, "_cleaned_ast_dirs", set()) + r = extract(paths, root=tmp_path, cache_root=tmp_path, parallel=False) + lbl = {n["id"]: n["label"] for n in r["nodes"]} + calls = {(lbl.get(e["source"]), lbl.get(e["target"])) for e in r["edges"] if e["relation"] == "calls"} + assert ("run()", ".send()") in calls + + +def test_warm_cache_matches_cold(tmp_path): + files = { + "client.py": _CLIENT, + "use.py": "from .client import Client\n\ndef run(c: Client):\n return c.send()\n", + } + _, cold, _ = _calls(tmp_path, files) + _, warm, _ = _calls(tmp_path, files) + key = lambda e: (e["source"], e["target"], e["confidence"]) # noqa: E731 + assert sorted(map(key, cold)) == sorted(map(key, warm)) + assert any(e["source"].endswith("_run") for e in warm)