Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 17 additions & 1 deletion graphtage/pydiff.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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:
Expand Down
41 changes: 39 additions & 2 deletions test/test_pydiff.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand All @@ -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"))
Expand Down