From 9c9445c1f31a09cd968fd894b37da2eefb39d338 Mon Sep 17 00:00:00 2001 From: Evan Sultanik Date: Wed, 9 Sep 2026 10:17:01 -0400 Subject: [PATCH] Build a Graphtage tree from plain import statements ASTBuilder registered an expander and a builder for ast.ImportFrom but nothing for ast.Import, so ast_to_tree raised NotImplementedError on any module containing a plain `import x` statement. Because `import x` is far more common than `from y import x`, that left most real Python sources undiffable. Register ast.Import on the existing alias expander, which both node types expose as `.names`, and add a builder that constructs the same graphtage.ast.Import node with an empty from_name. Both the node and PyImportFormatter already treat an empty from_name as the plain form, so no change to the printing side is needed. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01GypKU5KdLfs2Cf8kS2TzJa --- graphtage/pydiff.py | 18 +++++++++++++++++- test/test_pydiff.py | 42 +++++++++++++++++++++++++++++++++++++++++- 2 files changed, 58 insertions(+), 2 deletions(-) diff --git a/graphtage/pydiff.py b/graphtage/pydiff.py index a68313c..49ce9a0 100644 --- a/graphtage/pydiff.py +++ b/graphtage/pydiff.py @@ -204,8 +204,9 @@ def build_call(self, _, children: list[TreeNode]): CallKeywords(()) ) + @Builder.expander(ast.Import) @Builder.expander(ast.ImportFrom) - def expand_import_from(self, node: ast.ImportFrom): + def expand_import(self, node: ast.Import | ast.ImportFrom): return node.names @Builder.builder(ast.ImportFrom) @@ -216,6 +217,21 @@ def build_import_from(self, node: ast.ImportFrom, children: list[TreeNode]): from_name = StringNode(node.module, quoted=False) return Import(names=ListNode(children), from_name=from_name) + @Builder.builder(ast.Import) + def build_import(self, _, children: list[TreeNode]): + """Builds a plain ``import x`` statement. + + :class:`ast.Import` has no module of its own, so the resulting :class:`graphtage.ast.Import` gets an empty + ``from_name``, which is how both the node and its formatter distinguish ``import x`` from ``from y import x``. + + Args: + children: The :class:`PyAlias` nodes built from the statement's aliases. + + Returns: + Import: The resulting node. + """ + return Import(names=ListNode(children), from_name=StringNode("", quoted=False)) + @Builder.builder(ast.alias) def build_alias(self, node: ast.alias, _): if not node.asname: diff --git a/test/test_pydiff.py b/test/test_pydiff.py index ed70fd8..bbf0352 100644 --- a/test/test_pydiff.py +++ b/test/test_pydiff.py @@ -1,13 +1,23 @@ import ast import dataclasses +from io import StringIO from unittest import TestCase import graphtage -from graphtage.pydiff import PyDiffFormatter, ast_to_tree, build_tree, print_diff +from graphtage.ast import Import +from graphtage.pydiff import PyAlias, PyDiffFormatter, ast_to_tree, build_tree, print_diff from .timing import run_with_time_limit +def format_tree(tree: graphtage.TreeNode) -> str: + stream = StringIO() + printer = graphtage.printer.Printer(out_stream=stream, ansi_color=False) + with printer: + PyDiffFormatter.DEFAULT_INSTANCE.print(printer, tree) + return stream.getvalue() + + class TestPyDiff(TestCase): def test_build_tree(self): self.assertIsInstance(build_tree([1, 2, 3, 4]), graphtage.ListNode) @@ -22,6 +32,36 @@ def test_python_list_literals_stay_ordered(self): for node in tree.dfs(): self.assertNotIsInstance(node, graphtage.UnorderedListNode) + def test_plain_import_builds(self): + """Reproduces https://github.com/trailofbits/graphtage/issues/151. + + `ASTBuilder` had no builder for `ast.Import`, so every module containing a plain `import` statement raised + `NotImplementedError`. + + """ + expected = { + "import os": [("os", "")], + "import os.path": [("os.path", "")], + "import os as o": [("os", "o")], + "import os, sys": [("os", ""), ("sys", "")], + } + for source, aliases in expected.items(): + with self.subTest(source=source): + imports = [node for node in ast_to_tree(ast.parse(source)).dfs() if isinstance(node, Import)] + self.assertEqual(1, len(imports)) + self.assertEqual("", imports[0].from_name.object) + names = imports[0].names.children() + for name, (expected_name, expected_as_name) in zip(names, aliases, strict=True): + self.assertIsInstance(name, PyAlias) + self.assertEqual(expected_name, name.name.object) + self.assertEqual(expected_as_name, name.as_name.object) + + def test_plain_import_printing(self): + """A plain `import` must not render the `from` clause that `from x import y` gets.""" + for source in ("import os", "import os.path", "import os, sys", "from os import path"): + with self.subTest(source=source): + self.assertEqual(f"{source}\n", format_tree(ast_to_tree(ast.parse(source)))) + def test_diff(self): t1 = [1, 2, {3: "three"}, 4] t2 = [1, 2, {3: 3}, "four"]