如何将 python 比较 ast 节点转换为 c?

BPL*_*BPL 7 c python compiler-construction abstract-syntax-tree transpiler

让我们从考虑python3.8.5 的语法开始,在这种情况下,我有兴趣弄清楚如何将 python 比较转换为 c。

为了简单起见,我们假设我们正在处理一个非常小的 Python 平凡子集,我们只想转译平凡的 Compare 表达式:

expr = Compare(expr left, cmpop* ops, expr* comparators)
Run Code Online (Sandbox Code Playgroud)

如果我没记错的话,在 python 中,诸如这样的表达式a<b<c被转换成类似a<b && b<cwhere b 只被评估一次的东西......所以我想在 c 中你应该做类似的事情bool v0=a<b; bool v1=v0<c,以防止 b 被多次评估,以防第一条是真的。

不幸的是,我不知道如何将其放入代码中,到目前为止,这就是我所拥有的:

import ast
import shutil
import textwrap
from subprocess import PIPE
from subprocess import Popen


class Visitor(ast.NodeVisitor):
    def visit(self, node):
        ret = super().visit(node)
        if ret is None:
            raise Exception("Unsupported node")
        return ret

    def visit_Expr(self, node):
        return f"{self.visit(node.value)};"

    def visit_Eq(self, node):
        return "=="

    def visit_Lt(self, node):
        return "<"

    def visit_LtE(self, node):
        return "<="

    def visit_Load(self, node):
        return "//load"

    def visit_Name(self, node):
        return f"{node.id}"

    def visit_Compare(self, node):
        left = self.visit(node.left)
        ops = [self.visit(x) for x in node.ops]
        comparators = [self.visit(x) for x in node.comparators]

        if len(ops) == 1 and len(comparators) == 1:
            return f"({left} {ops[0]} {comparators[0]})"
        else:
            lhs = ",".join([f"'{v}'" for v in ops])
            rhs = ",".join([f"{v}" for v in comparators])
            return f"cmp<{lhs}>({rhs})"

    def visit_Call(self, node):
        func = self.visit(node.func)
        args = [self.visit(x) for x in node.args]
        # keywords = [self.visit(x) for x in node.keywords]
        return f"{func}({','.join(args)})"

    def visit_Module(self, node):
        return f"{''.join([self.visit(x) for x in node.body])}"

    def visit_Num(self, node):
        return node.n


if __name__ == "__main__":
    out = Visitor().visit(
        ast.parse(
            textwrap.dedent(
                """
            1 == 1<3
            1 == (1<3)
            1 == (0 < foo(0 <= bar() < 3, baz())) < (4 < 5)
            foo(0 <= bar() < 3, baz())
        """
            )
        )
    )

    if shutil.which("clang-format"):
        cmd = "clang-format -style webkit -offset 0 -length {} -assume-filename None"
        p = Popen(
            cmd.format(len(out)), stdout=PIPE, stdin=PIPE, stderr=PIPE, shell=True
        )
        out = p.communicate(input=out.encode("utf-8"))[0].decode("utf-8")
        print(out)
    else:
        print(out)
Run Code Online (Sandbox Code Playgroud)

如您所见,输出将是某种不可编译的 c 输出:

cmp<'==', '<'>(1, 3);
(1 == (1 < 3));
cmp<'==', '<'>((0 < foo(cmp<'<=', '<'>(bar(), 3), baz())), (4 < 5));
foo(cmp<'<=', '<'>(bar(), 3), baz());
Run Code Online (Sandbox Code Playgroud)

问题,什么是算法(python 工作示例在这里是理想的,但只有一些允许我改进提供的片段的通用伪代码也可以)允许我将 python 比较表达式转换为 c?

Ste*_*cht 2

转换 Compare 表达式时的另一个复杂问题是,您希望防止在拆分后多次使用的子表达式被多次求值,如果存在函数调用等副作用,这一点尤其重要。

人们可以提前将子表达式声明为变量,以避免多次求值。

Alexander Schepanovski 提出了一种将 Python 比较表达式转换为 JavaScript 的巧妙方法。他在博客文章中详细解释了他的整个解决方案:http://hackflow.com/blog/2015/04/12/metaprogramming-beyond-decency-part-2/。

基本上同样可以应用于 C 的转译。

他确定相邻操作数对。这对于将链式比较转换为单独的比较是必要的,其中“中间”操作数随后被复制,并且是分割的第二个子比较的左操作数。

可以使用一种符号表将变量与子表达式相关联。变量的命名可以通过一个简单的计数器来完成。

访问表达式节点时可以输出变量。要获得问题中作为示例给出的表达式的 C 输出,您可以简单地发出 printf。

为了进一步简化,我们可以假设假设的小而平凡的 Python 子集仅处理 int 表达式。

Python代码

我已经获取了您的代码片段,并根据上述几点对其进行了稍微修改,使其成为一个独立的示例,可为您的示例表达式输出可编译的 C 代码。

import ast
import itertools
import textwrap


def pairwise(iterable):
    """s -> (s0,s1), (s1,s2), (s2, s3), ..."""
    a, b = itertools.tee(iterable)
    next(b, None)
    return zip(a, b)


class Visitor(ast.NodeVisitor):
    def __init__(self):
        self.varCounter = 0
        self.varTable = []

    def visit_Expr(self, node):
        code = self.visit(node.value)
        variables = '\n'.join(self.varTable)
        self.varTable = []
        return f'{variables}\nprintf("%d\\n", {code});\n'

    def visit_Eq(self, node):
        return "=="

    def visit_Lt(self, node):
        return '<'

    def visit_LtE(self, node):
        return '<='

    def visit_Gt(self, node):
        return ">"

    def visit_GtE(self, node):
        return ">="

    def visit_Name(self, node):
        return str(node.id)

    # see http://hackflow.com/blog/2015/04/12/metaprogramming-beyond-decency-part-2/
    def visit_Compare(self, node):
        ops = node.ops
        operands = [node.left] + node.comparators
        variables = []
        for o in operands:
            self.varCounter += 1
            num = self.varCounter
            op = self.visit(o)
            variables.append((num, op))
            self.varTable.append(f'int t{num} = {op};')

        pairs = pairwise(variables)  # adjacent pairs of operands

        return ' && '.join('%s(%s %s %s)' %
                             ('!' if isinstance(op, ast.NotIn) else '',
                              f't{l[0]}', self.visit(op), f't{r[0]}')
                             for op, (l, r) in zip(ops, pairs))

    def visit_Call(self, node):
        args = [self.visit(x) for x in node.args]
        return self.visit(node.func) + "(" + ", ".join(args) + ")"

    def visit_Num(self, node):
        return str(node.n)


def main():
    analyzer = Visitor()
    tree = ast.parse(
        textwrap.dedent(
            """
            1 == 1<3
            1 == (1<3)
            1 == (0 < foo(0 <= bar() < 3, baz())) < (4 < 5)
            foo(0 <= bar() < 3, baz())
            """
        )
    )

    # print(ast.dump(tree))

    for node in ast.iter_child_nodes(tree):
        c = analyzer.visit(node)
        print(c)


if __name__ == '__main__':
    main()
Run Code Online (Sandbox Code Playgroud)

测试运行

运行Python程序时,调试控制台中会显示以下内容:

int t1 = 1;
int t2 = 1;
int t3 = 3;
printf("%d\n", (t1 == t2) && (t2 < t3));

int t4 = 1;
int t6 = 1;
int t7 = 3;
int t5 = (t6 < t7);
printf("%d\n", (t4 == t5));

int t8 = 1;
int t10 = 0;
int t12 = 0;
int t13 = bar();
int t14 = 3;
int t11 = foo((t12 <= t13) && (t13 < t14), baz());
int t9 = (t10 < t11);
int t16 = 4;
int t17 = 5;
int t15 = (t16 < t17);
printf("%d\n", (t8 == t9) && (t9 < t15));

int t18 = 0;
int t19 = bar();
int t20 = 3;
printf("%d\n", foo((t18 <= t19) && (t19 < t20), baz()));
Run Code Online (Sandbox Code Playgroud)

当然,有一种方法可以进一步简化这一点。例如,常量表达式不需要分配给变量。当然,还有更多细节需要考虑。但这应该是为示例数据输出可编译 C 代码的起点。