《Python高级编程》2.3 动态代码生成与 AST 变换

本节从字符串拼接式的代码生成升级到 AST 变换:用 compile() 的三种模式、ast.NodeTransformer 与 fix_missing_locations 把 assert 改写成显式检查,并注入函数计时。实测改写前后的 dis 字节码差异,以及为什么改写版能在 python -O 下继续生效,而原始 assert 会被剥离。

本节目标:掌握用 ast 模块解析、改写、重新编译代码的完整链路,并能用它实现有实际意义的代码变换。
适用版本:Python 3.12+(实测 3.14.6)

2.3 动态代码生成与 AST 变换

站内专题 Python 元编程与动态特性深度解析 在讲代码生成时,用的是字符串拼接 + exec——把 f"self.{field} = {field}" 拼成源码字符串再执行。那是最容易想到的做法,但也最容易出错:拼字符串没有语法检查,字段名里混进一个引号或换行,得到的就是一段半损坏的代码。专题结尾也留了一句「生产环境应优先使用 ast 模块」,但没展开。这一节就把那条路走完:解析成语法树、在树层面改写、再编译回代码对象,全程不碰字符串拼接。

2.3.1 从源码到代码对象的三个入口

Python 把源码变成可执行对象,有一条清晰的流水线:源码字符串 → ast.parse → AST → compile → code 对象。三个入口各有分工:

入口输入输出典型用途
compile(src, file, mode)字符串code 对象需要控制编译模式时
eval(code)表达式值只求值一个表达式
exec(code)语句块None执行多条语句

compile 的 mode 有三种,实测一下:

print("eval  :", eval(compile("1 + 2", "<s>", "eval")))
print("exec  :", compile("x = 1", "<s>", "exec").co_name)
print("single:", compile("print(1)", "<s>", "single").co_name)

真实输出:

eval  : 3
exec  : <module>
single: <module>
  • "eval":只能编译单个表达式,返回表达式的值。
  • "exec":编译一整段语句,适合模块级代码。
  • "single":编译交互式单条语句,会额外生成 PRINT_EXPR 之类的行为,一般只在 REPL 里用。

ast.parse 等价于 compile(src, "<unknown>", "exec", ast.PyCF_ONLY_AST),返回的是语法树而不是 code 对象。要改写代码,就必须停在这一步。

2.3.2 AST 的节点不变式

AST 不是随便嵌套的对象,它有两类容易踩的约束。

约束一:每个语句节点必须带位置信息(lineno / col_offset)。 编译器靠它生成回溯信息。手动构造节点时如果忘了填,编译直接失败:

import ast

tree = ast.parse("x = 1\n")
tree.body.append(ast.Expr(value=ast.Constant(value="hello")))  # 新节点没有位置
try:
    compile(tree, "<no-loc>", "exec")
except TypeError as e:
    print("缺少位置信息:", e)
ast.fix_missing_locations(tree)
print("修复后可编译:", compile(tree, "<fixed>", "exec").co_name)

真实输出:

缺少位置信息: required field "lineno" missing from stmt
修复后可编译: <module>

ast.fix_missing_locations() 是变换器的必备收尾步骤——它把缺失的位置信息从父节点继承下来。绝大多数「AST 变换后编译报 lineno missing」的报错,都是漏了这一句。

约束二:名字节点带 ctx(上下文)。 同一个 x,在赋值左边是 Store,在右边是 Load:

import ast
print(ast.dump(ast.parse("x = x + 1").body[0], indent=2))

真实输出:

Assign(
  targets=[
    Name(id='x', ctx=Store())],
  value=BinOp(
    left=Name(id='x', ctx=Load()),
    op=Add(),
    right=Constant(value=1)))

写变换器时构造 Name 节点,必须显式给出 ctx:读取用 ast.Load(),赋值用 ast.Store()。给错方向不会报语法错,但会生成行为错误的字节码。

2.3.3 实战一:把 assert 改写成显式检查

先看原始 assert 编译成什么样:

import ast, dis

src = '''
def area(w, h):
    assert w > 0 and h > 0, "尺寸必须为正"
    return w * h
'''
ns = {}
exec(compile(src, "<orig>", "exec"), ns)
dis.dis(ns['area'])

真实输出:

  3           LOAD_FAST_BORROW         0 (w)
              LOAD_SMALL_INT           0
              COMPARE_OP             148 (bool(>))
              POP_JUMP_IF_FALSE        8 (to L1)
              NOT_TAKEN
              LOAD_FAST_BORROW         1 (h)
              LOAD_SMALL_INT           0
              COMPARE_OP             148 (bool(>))
              POP_JUMP_IF_TRUE         8 (to L2)
              NOT_TAKEN
      L1:     LOAD_COMMON_CONSTANT     0 (AssertionError)
              LOAD_CONST               1 ('尺寸必须为正')
              CALL                     0
              RAISE_VARARGS            1

条件不成立时,它加载 AssertionError(3.14 用了 LOAD_COMMON_CONSTANT 这个专用指令)并抛出。现在写一个 NodeTransformer,把 assert test, msg 改写成 if not test: raise AssertionError(msg):

import ast

class AssertToCheck(ast.NodeTransformer):
    def visit_Assert(self, node):
        test = ast.UnaryOp(op=ast.Not(), operand=node.test)
        msg = node.msg if node.msg is not None else ast.Constant(value="assertion failed")
        raise_stmt = ast.Raise(
            exc=ast.Call(func=ast.Name(id="AssertionError", ctx=ast.Load()),
                         args=[msg], keywords=[]),
            cause=None)
        new = ast.If(test=test, body=[raise_stmt], orelse=[])
        return ast.copy_location(new, node)   # 继承原 assert 的位置

tree = ast.parse(src)
tree = AssertToCheck().visit(tree)
ast.fix_missing_locations(tree)
print(ast.unparse(tree))

ast.unparse 把改写后的语法树反编译回源码,这是验证变换结果最方便的手段。真实输出:

def area(w, h):
    if not (w > 0 and h > 0):
        raise AssertionError('尺寸必须为正')
    return w * h

assert 已经消失,变成了显式的 if ... raise。注意 ast.copy_location:把新 If 节点的位置对齐到原 assert,这样运行时报错的行号依然指向源码里写 assert 的那一行。

2.3.4 改写前后的 dis 对比

把改写后的树编译出来,再 dis 一次:

ns2 = {}
exec(compile(tree, "<rewritten>", "exec"), ns2)
dis.dis(ns2['area'])

真实输出:

  3           LOAD_FAST_BORROW         0 (w)
              LOAD_SMALL_INT           0
              COMPARE_OP             148 (bool(>))
              POP_JUMP_IF_FALSE        8 (to L1)
              NOT_TAKEN
              LOAD_FAST_BORROW         1 (h)
              LOAD_SMALL_INT           0
              COMPARE_OP             148 (bool(>))
              POP_JUMP_IF_TRUE        12 (to L2)
              NOT_TAKEN
      L1:     LOAD_GLOBAL              1 (AssertionError + NULL)
              LOAD_CONST               1 ('尺寸必须为正')
              CALL                     1
              RAISE_VARARGS            1

条件判断部分(前 10 条指令)完全一致——因为 assert 的条件本来就是这样求值的。差异只在抛错那三条:

指令原始 assert改写后
加载异常类LOAD_COMMON_CONSTANT AssertionErrorLOAD_GLOBAL AssertionError + NULL
构造异常CALL 0CALL 1(消息作为参数)
抛出RAISE_VARARGS 1RAISE_VARARGS 1

原始 assert 用 LOAD_COMMON_CONSTANT(3.14 为高频常量做的专用指令)+ CALL 0,比改写版的 LOAD_GLOBAL + CALL 1 略快,但差别只在断言失败那条冷路径上。成功路径两边开销相同。

2.3.5 为什么值得改写:-O 会剥掉 assert

如果改写前后只差冷路径的几条指令,何必大费周章?真正的理由在 python -O。优化模式下 CPython 会直接删除所有 assert 语句——它们连字节码都不生成。用同一段源码在默认和 -O 两种模式下各跑一次:

=== 默认 ===
sys.flags.optimize = 0
assert 版 area(-1,4): AssertionError: 尺寸必须为正
改写版 area(-1,4): AssertionError: 尺寸必须为正
=== -O ===
sys.flags.optimize = 1
assert 版 area(-1,4): -4
改写版 area(-1,4): AssertionError: 尺寸必须为正

在 -O 下:

  • assert 版:断言被剥离,area(-1, 4) 直接算出 -1 * 4 = -4 返回——非法参数悄无声息地通过了。
  • 改写版:if not ...: raise 是普通控制流,不受 -O 影响,依然抛 AssertionError。

这就是 AST 变换的实际价值:把「调试期才生效的检查」变成「任何模式下都生效的检查」,而源码可以继续写更简洁的 assert。同理,把「输入校验」从 assert 改写成显式 raise,正是很多框架在导入钩子里做的事。

2.3.6 实战二:注入函数计时

第二个例子更有工程味:给每个函数体自动套一层计时。变换器在每个函数体开头插入 __t0 = time.perf_counter(),并用 try/finally 在退出时打印耗时:

import ast

class InjectTiming(ast.NodeTransformer):
    def visit_FunctionDef(self, node):
        self.generic_visit(node)
        name = node.name
        t0 = ast.Assign([ast.Name("__t0", ast.Store())],
            ast.Call(ast.Attribute(ast.Name("time", ast.Load()), "perf_counter", ast.Load()), [], []))
        report = ast.Expr(ast.Call(ast.Name("print", ast.Load()),
            [ast.JoinedStr([ast.Constant(f"[{name}] elapsed "),
                ast.FormattedValue(ast.BinOp(
                    ast.Call(ast.Attribute(ast.Name("time", ast.Load()), "perf_counter", ast.Load()), [], []),
                    ast.Sub(), ast.Name("__t0", ast.Load())), -1, ast.Constant(".6f"))])], []))
        node.body = [t0, ast.Try(node.body, [], [], [report])]
        return ast.fix_missing_locations(node)

generic_visit(node) 先递归处理子节点,再改写当前节点——顺序不能反,否则嵌套函数会被处理两次。跑一段真实负载:

src = '''
import time
def total(n):
    s = 0
    for i in range(n):
        s += i * i
    return s
'''
tree = InjectTiming().visit(ast.parse(src))
ast.fix_missing_locations(tree)
ns = {}
exec(compile(tree, "<timed>", "exec"), ns)
print("total(1_000_000) =", ns['total'](1_000_000))

真实输出:

[total] elapsed 0.047563
total(1_000_000) = 333332833333500000

try/finally 是关键:无论函数是正常 return 还是抛异常退出,finally 都会执行,计时都不会漏。编译后的字节码里能看到这段结构(节选):

   4           CALL                     0
               STORE_FAST               1 (__t0)
   8           LOAD_FAST_BORROW         2 (s)
   4   L4:     LOAD_GLOBAL              7 (print + NULL)
               ...
               RETURN_VALUE
  --   L5:     PUSH_EXC_INFO
   4           LOAD_GLOBAL              7 (print + NULL)

L4 是正常返回路径上的计时打印,L5 是异常路径——同一个 finally 块被编译器复制成了两份,这就是 try/finally 在字节码层面的样子。

2.3.7 工程边界

AST 变换很强,但有明确的适用边界,别把它当成万能工具:

  • 它发生在编译期,不是运行期。 上面所有例子都是「拿到源码字符串 → 变换 → 编译」。要让变换作用于 import 进来的模块,得写导入钩子在加载器里插入变换步骤——那是第 7 章的内容。
  • ast.unparse 不是无损往返。 它能还原语义,但注释、空行、引号风格都会丢失。别用它做「格式化」。
  • AST 节点随版本演进。 ast 的字段在 3.8~3.14 间多次增删,跨版本运行的变换器必须用 getattr / hasattr 做兼容,不要硬编码字段名。
  • 优先用标准库已经封装的变换。 想让断言永远生效,优先用 pytest 的断言重写;想做类型检查,用 mypy。手写 NodeTransformer 只在没有现成工具、且变换规则确实简单时才是对的。

小结

  • 代码生成的正路是 ast.parse → NodeTransformer → ast.fix_missing_locations → compile,而不是拼字符串 + exec。
  • compile 有三种模式:eval(表达式)、exec(语句块)、single(交互式)。
  • AST 有两条硬约束:语句节点必须带位置信息(靠 fix_missing_locations 补齐),名字节点必须给对 ctx。
  • 把 assert 改写成 if not ...: raise 后,条件判断的字节码不变,差别只在抛错冷路径;但改写版能在 python -O 下继续生效。
  • 注入计时的变换展示了 generic_visit 先递归、try/finally 保证异常路径也计时的要点。
  • AST 变换是编译期行为,要作用于导入模块需配合导入钩子;ast.unparse 不保证源码级往返。

我们已经能在语法树层面改写代码了,但改写的产物最终还是会变成字节码。下一节就从「语法树」下降到「指令」——看 CPython 的虚拟机到底怎么执行这些字节码。

阅读导航:上一节:2.2 类装饰器、set_name 与属性工厂 · 下一节:3.1 CPython 执行模型与字节码 。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「python」更多文章

  1. 《Python高级编程》目录
  2. 《Python高级编程》11.3 PEP 流程与版本迁移策略
  3. 《Python高级编程》11.2 嵌入式与自由线程运行时