How can I get the next sibling of a node in a Python AST?

Viewed 93

I write flake8-simplify (PyPI, GitHub), a Flake8 plugin that finds code which can be simplified.

One pattern I want to find is this:

# Bad
for x in iterable:
    if check(x):
        return True
return False

# Good
return any(check(x) for x in iterable)

The AST of the "bad" part looks like this:

[
        For(
            target=Name(id='x', ctx=Store()),
            iter=Name(id='iterable', ctx=Load()),
            body=[
                If(
                    test=Call(
                        func=Name(id='check', ctx=Load()),
                        args=[Name(id='x', ctx=Load())],
                        keywords=[],
                    ),
                    body=[
                        Return(
                            value=Constant(value=True, kind=None),
                        ),
                    ],
                    orelse=[],
                ),
            ],
            orelse=[],
            type_comment=None,
        ),
        Return(
            value=Constant(value=False, kind=None),
        ),
    ]

The problem here is that I need to check the For and the Return. **How can I get the next (sibling) node, given a Python ast.For node?

What I tried

import ast
from typing import List, Tuple

SIM110 = "SIM110 Use 'return any({check} for {target} in {iterable})'"


class Visitor(ast.NodeVisitor):
    def visit_For(self, node: ast.For) -> None:
        self.errors += _get_sim110(node)
        self.generic_visit(node)


def _get_sim110(node: ast.ForOp) -> List[Tuple[int, int, str]]:
    errors: List[Tuple[int, int, str]] = []
    if not (
        len(node.body) == 1
        and isinstance(node.body[0], ast.If)
        and len(node.body[0].body) == 1
        and isinstance(node.body[0].body[0], ast.Return)
        and node.body[0].body[0].value is True
        and [TODO: after the For loop is a "return False"]
    ):
        return errors

    # Prepare the message
    check = to_source(node.body[0].test)
    target = to_source(node.target)
    iterable = to_source(node.iter)
    errors.append(
        (
            node.lineno,
            node.col_offset,
            SIM110.format(check=check, target=target, iterable=iterable),
        )
    )
    return errors


def to_source(node: ast.expr) -> str:
    import astor
    source = astor.to_source(node).strip()
    return source
0 Answers
Related