fix(extract): emit nodes for exported scalar bindings

This commit is contained in:
oleksii-tumanov
2026-07-30 17:13:48 +01:00
committed by safishamsi
parent bdf62ce275
commit d7cb5b1b1c
2 changed files with 113 additions and 6 deletions
+17 -6
View File
@@ -1773,9 +1773,10 @@ def _js_extra_walk(node, source: bytes, file_nid: str, stem: str, str_path: str,
# phantom god-nodes. Bodies of arrow functions are walked separately
# via function_bodies, so we never need to emit nodes for locals here.
parent = node.parent
is_exported = parent is not None and parent.type == "export_statement"
is_module_level = parent is not None and (
parent.type == "program"
or (parent.type == "export_statement"
or (is_exported
and parent.parent is not None
and parent.parent.type == "program")
)
@@ -1787,9 +1788,15 @@ def _js_extra_walk(node, source: bytes, file_nid: str, stem: str, str_path: str,
for child in node.children:
if child.type == "variable_declarator":
value = child.child_by_field_name("value")
name_node = child.child_by_field_name("name")
is_exported_scalar_binding = (
is_exported
and name_node is not None
and name_node.type == "identifier"
and bool(normalize_id(_read_text(name_node, source)))
)
if value and value.type in _JS_FUNCTION_VALUE_TYPES:
# `const f = () => {}` and `const f = function(){}`
name_node = child.child_by_field_name("name")
if name_node:
func_name = _read_text(name_node, source)
line = child.start_point[0] + 1
@@ -1809,11 +1816,15 @@ def _js_extra_walk(node, source: bytes, file_nid: str, stem: str, str_path: str,
if body:
function_bodies.append((func_nid, body))
arrow_found = True
elif value and value.type in (
"object", "array", "as_expression", "call_expression", "new_expression",
elif value and (
is_exported_scalar_binding
or value.type in (
"object", "array", "as_expression", "call_expression",
"new_expression",
)
):
# Module-level const with literal/object/array/factory value
name_node = child.child_by_field_name("name")
# Simple exported identifiers are part of the module API
# regardless of initializer shape. Keep other scalar noise suppressed.
if name_node:
const_name = _read_text(name_node, source)
line = child.start_point[0] + 1
+96
View File
@@ -0,0 +1,96 @@
from __future__ import annotations
import pytest
from graphify.extract import extract, extract_js
@pytest.mark.parametrize("suffix", [".js", ".ts"])
def test_exported_scalar_bindings_emit_nodes(tmp_path, suffix):
source = tmp_path / f"constants{suffix}"
source.write_text(
"""
export const NUMBER = 42;
export const STRING = "value";
export const BOOLEAN = true;
export const TEMPLATE = `value-${NUMBER}`;
export const MEMBER = process.env.VALUE;
export const LOGICAL = process.env.VALUE ?? "fallback";
export const TERNARY = BOOLEAN ? "yes" : "no";
const internalScalar = 1;
function helper() {
const localScalar = 2;
}
""",
encoding="utf-8",
)
result = extract_js(source)
labels = {node["label"] for node in result["nodes"]}
assert {
"NUMBER",
"STRING",
"BOOLEAN",
"TEMPLATE",
"MEMBER",
"LOGICAL",
"TERNARY",
} <= labels
assert "internalScalar" not in labels
assert "localScalar" not in labels
def test_exported_scalar_fix_skips_unsupported_binding_patterns(tmp_path):
source = tmp_path / "patterns.ts"
source.write_text(
"""
const config = { source: 1 };
const items = [1];
export const { source: renamed } = config;
export const [first] = items;
export const $ = 1;
export const _ = 2;
""",
encoding="utf-8",
)
result = extract_js(source)
labels = {node["label"] for node in result["nodes"]}
assert "$" not in labels
assert "_" not in labels
assert not any("renamed" in label or "first" in label for label in labels)
assert all(edge["source"] != edge["target"] for edge in result["edges"])
def test_exported_scalar_binding_satisfies_named_import_target(tmp_path):
exporter = tmp_path / "constants.ts"
exporter.write_text(
"""
export const A_PREFIX = process.env.A_PREFIX ?? "X>";
export const A_MAX = Number(process.env.A_MAX || 10);
""",
encoding="utf-8",
)
importer = tmp_path / "consumer.ts"
importer.write_text(
'import { A_PREFIX, A_MAX } from "./constants";\n',
encoding="utf-8",
)
result = extract(
[exporter, importer],
cache_root=tmp_path,
parallel=False,
)
node_ids = {node["id"] for node in result["nodes"]}
import_targets = {
edge["target"]
for edge in result["edges"]
if edge["relation"] == "imports"
}
assert import_targets
assert import_targets <= node_ids