本节目标:掌握用
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 AssertionError | LOAD_GLOBAL AssertionError + NULL |
| 构造异常 | CALL 0 | CALL 1(消息作为参数) |
| 抛出 | RAISE_VARARGS 1 | RAISE_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 执行模型与字节码 。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。