diff --git a/graphtage/pydiff.py b/graphtage/pydiff.py index b92d8f5..6db44dc 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 effb501..da01836 100644 --- a/test/test_pydiff.py +++ b/test/test_pydiff.py @@ -4,13 +4,21 @@ from unittest import TestCase import graphtage -from graphtage.ast import Subscript +from graphtage.ast import Import, Subscript from graphtage.printer import Printer -from graphtage.pydiff import PyDiffFormatter, PyObjAttribute, ast_to_tree, build_tree, print_diff +from graphtage.pydiff import PyAlias, PyDiffFormatter, PyObjAttribute, ast_to_tree, build_tree, print_diff from .timing import run_with_time_limit +def format_tree(tree: graphtage.TreeNode) -> str: + stream = StringIO() + printer = Printer(out_stream=stream, ansi_color=False) + with printer: + PyDiffFormatter.DEFAULT_INSTANCE.print(printer, tree) + return stream.getvalue() + + def render_diff(from_source: str, to_source: str) -> str: """Diffs two Python sources and returns the rendering produced by :class:`PyDiffFormatter`.""" from_tree = ast_to_tree(ast.parse(from_source)) @@ -36,6 +44,35 @@ 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_attribute_receiver_is_not_quoted(self): stream = StringIO() attribute = PyObjAttribute(graphtage.StringNode("package"), graphtage.StringNode("member"))