Files
TransPyC/lib/core/DecoratorPass.py
2026-07-18 19:25:40 +08:00

548 lines
22 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
DecoratorPass - Viper 语言自定义装饰器的 IR 层包装转换 (v4)
v4 增强版特性:
1. 栈帧局部上下文 (ctx):替代全局变量,解决并发/递归状态错乱
2. 流程劫持handler 返回 i32 控制原函数调用次数
- 返回 0跳过原函数调用@cache, @mock
- 返回 1正常调用一次默认
- 返回 N > 1循环调用 N 次(@repeat
3. alwaysinline 属性:消除 wrapper 调用开销
4. 参数打包与修改pre-phase 后从 args_struct 读回参数,装饰器可修改入参
5. 返回值修改ret_ptr 可写,装饰器可在后置阶段修改返回值
6. 递归装饰保护:最外层 wrapper 检查全局标志位,递归调用时跳过所有装饰逻辑
装饰器调用约定 (v4)
i32 decor_name(i8* ctx, i8* func_name, i32 phase,
i8* args_ptr, i8* ret_ptr, i8* decor_args_ptr)
参数说明:
ctx : i8* - 栈帧局部上下文缓冲区32字节每次调用独立分配
前置阶段可写入状态,后置阶段可读取,生命周期绑定本次调用
多线程/递归完全隔离,替代全局变量
func_name : i8* - 原始函数名字符串
phase : i32 - 执行阶段0=前置, 1=后置)
args_ptr : i8* - 指向函数参数打包结构体(可读写,无参数时为 null
前置阶段可写入修改参数值wrapper 会从结构体读回修改后的参数
ret_ptr : i8* - 指向返回值(前置阶段为 nullvoid 函数始终为 null
后置阶段可读写,装饰器可修改返回值
decor_args_ptr : i8* - 指向装饰器参数打包结构体(无参数装饰器为 null
返回值(仅前置阶段有效,后置阶段忽略):
0 : 跳过原函数调用(@cache, @mock, @skip
1 : 调用原函数一次(默认行为,@log, @timing, @trace
N>1: 循环调用原函数 N 次(@repeat, @benchmark
编译器生成的 wrapper以 add(i32, i32) -> i32, @log 为例):
define i32 @__decor_wrap_add(i32 %a, i32 %b) alwaysinline {
entry:
; 递归保护:检查标志位
%rec_val = Load i8, i8* @__decor_rec___decor_wrap_add
%is_rec = icmp ne i8 %rec_val, 0
br i1 %is_rec, label %recursive_call, label %decorated_entry
recursive_call:
; 递归调用:跳过所有装饰逻辑,直接调用原函数
%rec_result = call i32 @add(i32 %a, i32 %b)
ret i32 %rec_result
decorated_entry:
; 设置递归保护标志
store i8 1, i8* @__decor_rec___decor_wrap_add
; 分配栈帧局部上下文
%ctx = alloca [32 x i8]
%ctx_ptr = bitcast [32 x i8]* %ctx to i8*
; 打包函数参数到结构体(装饰器可通过 args_ptr 修改)
%args = alloca {i32, i32}
store i32 %a, i32* getelementptr({i32, i32}* %args, i32 0, i32 0)
store i32 %b, i32* getelementptr({i32, i32}* %args, i32 0, i32 1)
%args_ptr = bitcast {i32, i32}* %args to i8*
; 分配返回值存储(零初始化)
%ret = alloca i32
store i32 0, i32* %ret
; 前置阶段handler 返回调用次数
%n = call i32 @log(i8* %ctx_ptr, i8* "add", i32 0,
i8* %args_ptr, i8* null, i8* null)
br label %loop.header
loop.header:
%i = phi i32 [0, %decorated_entry], [%i.next, %loop.body]
%cond = icmp slt i32 %i, %n
br i1 %cond, label %loop.body, label %post
loop.body:
; 从结构体读回参数(装饰器可能已修改)
%a.Loaded = Load i32, i32* getelementptr({i32, i32}* %args, i32 0, i32 0)
%b.Loaded = Load i32, i32* getelementptr({i32, i32}* %args, i32 0, i32 1)
%result = call i32 @add(i32 %a.Loaded, i32 %b.Loaded)
store i32 %result, i32* %ret
%i.next = add i32 %i, 1
br label %loop.header
post:
%ret_ptr = bitcast i32* %ret to i8*
call i32 @log(i8* %ctx_ptr, i8* "add", i32 1,
i8* %args_ptr, i8* %ret_ptr, i8* null)
; 清除递归保护标志
store i8 0, i8* @__decor_rec___decor_wrap_add
%final = Load i32, i32* %ret
ret i32 %final
}
链式装饰器(从下到上嵌套):
@log
@timing
def f(x) -> int:
等价于 log(timing(f))
执行顺序log_pre → timing_pre → f → timing_post → log_post
生成__decor_wrap_f (log) → __decor_wrap_f_timing (timing) → f
递归保护仅在最外层 wrapper 生效,递归调用直接跳到原始 f
"""
from __future__ import annotations
import llvmlite.ir as ir
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from lib.core.LlvmCodeGenerator import LlvmCodeGenerator
# 栈帧局部上下文缓冲区大小(字节)
_CTX_SIZE: int = 32
# 装饰器处理函数的 LLVM 签名i32 (i8*, i8*, i32, i8*, i8*, i8*)
_DECOR_HANDLER_PARAM_TYPES: ir.FunctionType | None = None # 延迟初始化
def _get_decor_handler_func_type() -> ir.FunctionType:
"""获取装饰器处理函数的 LLVM FunctionType: i32 (i8*, i8*, i32, i8*, i8*, i8*)"""
global _DECOR_HANDLER_PARAM_TYPES
if _DECOR_HANDLER_PARAM_TYPES is None:
i8_ptr: ir.PointerType = ir.IntType(8).as_pointer()
_DECOR_HANDLER_PARAM_TYPES = ir.FunctionType(
ir.IntType(32),
[i8_ptr, i8_ptr, ir.IntType(32), i8_ptr, i8_ptr, i8_ptr]
)
return _DECOR_HANDLER_PARAM_TYPES
def _find_or_declare_handler(module: ir.Module, handler_name: str) -> ir.Function:
"""在 module 中查找或声明装饰器处理函数"""
f: ir.Function
for f in module.functions:
if f.name == handler_name:
return f
func_type: ir.FunctionType = _get_decor_handler_func_type()
handler: ir.Function = ir.Function(module, func_type, name=handler_name)
return handler
def _make_string_constant(module: ir.Module, text: str, name_prefix: str = "decor.str") -> ir.Constant:
"""在 module 中创建一个全局字符串常量,返回 i8* 指针"""
counter: int = getattr(_make_string_constant, '_counter', 0) + 1
_make_string_constant._counter = counter
const_name: str = f"{name_prefix}.{counter}"
encoded: bytes = text.encode('utf-8') + b'\x00'
const_type: ir.ArrayType = ir.ArrayType(ir.IntType(8), len(encoded))
const_val: ir.Constant = ir.Constant(const_type, bytearray(encoded))
global_var: ir.GlobalVariable = ir.GlobalVariable(module, const_type, name=const_name)
global_var.global_constant = True
global_var.linkage = 'private'
global_var.initializer = const_val
return global_var.bitcast(ir.IntType(8).as_pointer())
def _make_decor_args_constant(module: ir.Module, decorator_info: dict) -> ir.Constant:
"""
为带参数的装饰器生成全局常量结构体,返回 i8* 指针。
无参数装饰器返回 null。
支持的参数类型:
- int → i32 / i64
- float → f64
- bool → i1
- str → i8*(全局字符串常量)
"""
deco_args: list = decorator_info.get('args', [])
deco_kwargs: dict = decorator_info.get('kwargs', {})
if not deco_args and not deco_kwargs:
return ir.Constant(ir.IntType(8).as_pointer(), None)
field_types: list = []
field_values: list = []
arg: int | str | float | bool
for arg in deco_args:
llvm_type: ir.Type | None
llvm_val: ir.Constant | None
llvm_type, llvm_val = _python_value_to_llvm(module, arg)
if llvm_type is None:
return ir.Constant(ir.IntType(8).as_pointer(), None)
field_types.append(llvm_type)
field_values.append(llvm_val)
kw: str
for kw in sorted(deco_kwargs.keys()):
llvm_type: ir.Type | None
llvm_val: ir.Constant | None
llvm_type, llvm_val = _python_value_to_llvm(module, deco_kwargs[kw])
if llvm_type is None:
return ir.Constant(ir.IntType(8).as_pointer(), None)
field_types.append(llvm_type)
field_values.append(llvm_val)
struct_type: ir.LiteralStructType = ir.LiteralStructType(field_types)
struct_const: ir.Constant = ir.Constant(struct_type, field_values)
counter: int = getattr(_make_decor_args_constant, '_counter', 0) + 1
_make_decor_args_constant._counter = counter
global_name: str = f"__decor_args.{counter}"
global_var: ir.GlobalVariable = ir.GlobalVariable(module, struct_type, name=global_name)
global_var.global_constant = True
global_var.linkage = 'private'
global_var.initializer = struct_const
return global_var.bitcast(ir.IntType(8).as_pointer())
def _python_value_to_llvm(module: ir.Module, value: int | str | float | bool) -> tuple[ir.Type | None, ir.Constant | None]:
"""将 Python 值转换为 (LLVM类型, LLVM常量) 元组。"""
if isinstance(value, bool):
return ir.IntType(1), ir.Constant(ir.IntType(1), int(value))
elif isinstance(value, int):
if -2**31 <= value < 2**31:
return ir.IntType(32), ir.Constant(ir.IntType(32), value)
else:
return ir.IntType(64), ir.Constant(ir.IntType(64), value)
elif isinstance(value, float):
return ir.DoubleType(), ir.Constant(ir.DoubleType(), value)
elif isinstance(value, str):
str_const = _make_string_constant(module, value, name_prefix="decor.arg")
return ir.IntType(8).as_pointer(), str_const
else:
return None, None
def _ZeroConstant(llvm_type: ir.Type) -> ir.Constant | None:
"""为给定 LLVM 类型创建零值常量"""
if isinstance(llvm_type, ir.IntType):
return ir.Constant(llvm_type, 0)
elif isinstance(llvm_type, (ir.FloatType, ir.DoubleType)):
return ir.Constant(llvm_type, 0.0)
elif isinstance(llvm_type, ir.PointerType):
return ir.Constant(llvm_type, None)
elif isinstance(llvm_type, ir.ArrayType):
elem_zero: ir.Constant | None = _ZeroConstant(llvm_type.element)
if elem_zero is None:
return None
return ir.Constant(llvm_type, [elem_zero] * llvm_type.count)
else:
return None
def _generate_single_wrapper(module: ir.Module, original_func: ir.Function, wrapper_name: str, func_name_str: str,
decorator_info: dict, is_export: bool = False,
is_outermost: bool = False, true_original_func: ir.Function | None = None) -> ir.Function:
"""
为单个装饰器生成一层 wrapper 函数。
v4 Wrapper 结构(最外层含递归保护):
entry → 递归保护检查(仅最外层)
recursive_call → 递归时直接调用原函数,跳过所有装饰(仅最外层)
decorated_entry → 设置递归标志,分配 ctx/args/ret调用 handler 前置
loop.header → phi 计数器,比较 i < n_calls
loop.body → 从 args_struct 读回参数(支持修改),调用原函数
post → 调用 handler 后置,清除递归标志,返回结果
"""
decor_name: str = decorator_info['name']
handler: ir.Function = _find_or_declare_handler(module, decor_name)
func_name_const: ir.Constant = _make_string_constant(module, func_name_str, name_prefix="decor.fn")
func_type: ir.FunctionType = original_func.ftype
return_type: ir.Type = func_type.return_type
param_types: list[ir.Type] = [p.type for p in original_func.args]
is_void: bool = isinstance(return_type, ir.VoidType)
wrapper_type: ir.FunctionType = ir.FunctionType(return_type, param_types)
wrapper: ir.Function = ir.Function(module, wrapper_type, name=wrapper_name)
# 添加 alwaysinline 属性,优化器完全内联消除调用开销
wrapper.attributes.add('alwaysinline')
i32: ir.IntType = ir.IntType(32)
i8: ir.IntType = ir.IntType(8)
i8_ptr: ir.PointerType = ir.IntType(8).as_pointer()
phase_pre: ir.Constant = ir.Constant(i32, 0)
phase_post: ir.Constant = ir.Constant(i32, 1)
null_ptr: ir.Constant = ir.Constant(i8_ptr, None)
zero_i32: ir.Constant = ir.Constant(i32, 0)
one_i32: ir.Constant = ir.Constant(i32, 1)
one_i8: ir.Constant = ir.Constant(i8, 1)
zero_i8: ir.Constant = ir.Constant(i8, 0)
# 递归保护:仅最外层 wrapper 生成
rec_flag: ir.GlobalVariable | None = None
if is_outermost and true_original_func is not None:
rec_flag_name: str = f"__decor_rec_{wrapper_name}"
# 检查是否已存在
rec_flag = None
g: ir.GlobalValue
for g in module.global_values:
if g.name == rec_flag_name:
rec_flag = g
break
if rec_flag is None:
rec_flag = ir.GlobalVariable(module, i8, name=rec_flag_name)
rec_flag.global_constant = False
rec_flag.linkage = 'internal'
rec_flag.initializer = zero_i8
# 创建基本块
entry_block: ir.Block = wrapper.append_basic_block("entry")
recursive_call_block: ir.Block | None
decorated_entry_block: ir.Block | None
if rec_flag is not None:
recursive_call_block = wrapper.append_basic_block("recursive_call")
decorated_entry_block = wrapper.append_basic_block("decorated_entry")
else:
recursive_call_block = None
decorated_entry_block = None
loop_header_block: ir.Block = wrapper.append_basic_block("loop.header")
loop_body_block: ir.Block = wrapper.append_basic_block("loop.body")
post_block: ir.Block = wrapper.append_basic_block("post")
# ==== Entry block ====
builder: ir.IRBuilder = ir.IRBuilder(entry_block)
loop_predecessor: ir.Block
if rec_flag is not None:
# 递归保护:检查标志位
rec_val: ir.LoadInstr = builder.load(rec_flag, name="rec_val")
is_recursive: ir.ICMPInstr = builder.icmp_signed('!=', rec_val, zero_i8)
builder.cbranch(is_recursive, recursive_call_block, decorated_entry_block)
# ==== Recursive call block ====
builder = ir.IRBuilder(recursive_call_block)
# 递归调用:直接调用真正的原函数,跳过所有装饰逻辑
rec_call_args: list[ir.Argument] = [arg for arg in wrapper.args]
rec_ret_val: ir.CallInstr = builder.call(true_original_func, rec_call_args)
if is_void:
builder.ret_void()
else:
builder.ret(rec_ret_val)
# ==== Decorated entry block ====
builder = ir.IRBuilder(decorated_entry_block)
# 设置递归保护标志
builder.store(one_i8, rec_flag)
# 记录 loop.header 的前置块为 decorated_entry_block
loop_predecessor = decorated_entry_block
else:
loop_predecessor = entry_block
# 1. 分配栈帧局部上下文(每次调用独立,线程安全,递归安全)
ctx_type: ir.ArrayType = ir.ArrayType(ir.IntType(8), _CTX_SIZE)
ctx_alloca: ir.AllocaInstr = builder.alloca(ctx_type, name="ctx")
ctx_ptr: ir.BitCastInstr = builder.bitcast(ctx_alloca, i8_ptr)
# 2. 打包函数参数到结构体(装饰器可通过 args_ptr 读写修改)
args_alloca: ir.AllocaInstr | None = None
args_struct_type: ir.LiteralStructType | None = None
if param_types:
args_struct_type = ir.LiteralStructType(param_types)
args_alloca = builder.alloca(args_struct_type, name="args")
zero: ir.Constant = ir.Constant(i32, 0)
i: int
arg: ir.Argument
for i, arg in enumerate(wrapper.args):
field_idx: ir.Constant = ir.Constant(i32, i)
gep: ir.GEPInstr = builder.gep(args_alloca, [zero, field_idx], inbounds=True)
builder.store(arg, gep)
args_ptr: ir.BitCastInstr | ir.Constant = builder.bitcast(args_alloca, i8_ptr)
else:
args_ptr = null_ptr
# 3. 分配返回值存储(零初始化,确保跳过原函数时返回安全默认值)
ret_alloca: ir.AllocaInstr | None
if not is_void:
ret_alloca = builder.alloca(return_type, name="ret")
zero_val: ir.Constant | None = _ZeroConstant(return_type)
if zero_val is not None:
builder.store(zero_val, ret_alloca)
else:
ret_alloca = None
# 4. 获取装饰器参数指针
decor_args_ptr: ir.Constant = _make_decor_args_constant(module, decorator_info)
# 5. 前置阶段handler 返回调用次数
n_calls: ir.CallInstr = builder.call(handler, [ctx_ptr, func_name_const, phase_pre,
args_ptr, null_ptr, decor_args_ptr])
# 6. 跳转到循环头
builder.branch(loop_header_block)
# ==== Loop header block ====
builder = ir.IRBuilder(loop_header_block)
i_phi: ir.PhiInstr = builder.phi(i32, "i")
i_phi.add_incoming(zero_i32, loop_predecessor)
cond: ir.ICMPInstr = builder.icmp_signed('<', i_phi, n_calls)
builder.cbranch(cond, loop_body_block, post_block)
# ==== Loop body block ====
builder = ir.IRBuilder(loop_body_block)
# v4: 从 args_struct 读回参数(装饰器可能在 pre-phase 中修改了参数)
call_args: list[ir.LoadInstr | ir.Argument] = []
if args_alloca is not None and args_struct_type is not None:
zero = ir.Constant(i32, 0)
for i in range(len(param_types)):
field_idx: ir.Constant = ir.Constant(i32, i)
gep: ir.GEPInstr = builder.gep(args_alloca, [zero, field_idx], inbounds=True)
Loaded_arg: ir.LoadInstr = builder.load(gep, name=f"arg.{i}")
call_args.append(Loaded_arg)
ret_val: ir.CallInstr = builder.call(original_func, call_args)
if ret_alloca is not None:
builder.store(ret_val, ret_alloca)
i_next: ir.AddInstr = builder.add(i_phi, one_i32)
i_phi.add_incoming(i_next, loop_body_block)
builder.branch(loop_header_block)
# ==== Post block ====
builder = ir.IRBuilder(post_block)
# 后置阶段
ret_ptr: ir.BitCastInstr | ir.Constant
if ret_alloca is not None:
ret_ptr = builder.bitcast(ret_alloca, i8_ptr)
else:
ret_ptr = null_ptr
builder.call(handler, [ctx_ptr, func_name_const, phase_post,
args_ptr, ret_ptr, decor_args_ptr])
# 递归保护:清除标志(必须在 post-phase 之后、return 之前)
if rec_flag is not None:
builder.store(zero_i8, rec_flag)
# 返回
if is_void:
builder.ret_void()
else:
final_ret: ir.LoadInstr = builder.load(ret_alloca, name="final_ret")
builder.ret(final_ret)
return wrapper
def _redirect_calls(module: ir.Module, old_func: ir.Function, new_func: ir.Function, skip_func_names: set[str] | None = None) -> None:
"""
替换 module 中所有对 old_func 的调用指令为对 new_func 的调用。
跳过 skip_func_names 中的函数wrapper 函数内部不应被重定向)。
"""
if skip_func_names is None:
skip_func_names = set()
func: ir.Function
for func in module.functions:
if func.name in skip_func_names:
continue
if func.name == new_func.name:
continue
block: ir.Block
for block in func.blocks:
instr: ir.Instruction
for instr in list(block.instructions):
if not isinstance(instr, ir.CallInstr):
continue
callee: ir.Function = instr.callee
if isinstance(callee, ir.Function) and callee is old_func:
instr.callee = new_func
def run(Gen: LlvmCodeGenerator) -> None:
"""
执行 DecoratorPass。
处理流程:
1. 遍历 _decorated_funcs 中所有带装饰标记的函数
2. 对每个函数,按装饰器列表从底到顶生成嵌套 wrapper
3. 将原函数改为 internal 链接
4. 替换所有调用点
"""
if not Gen._decorated_funcs:
return
module: ir.Module = Gen.module
decorated: dict = Gen._decorated_funcs.copy()
mangled_name: str
info: dict
for mangled_name, info in decorated.items():
decorators: list = info['decorators']
func_name: str = info['func_name']
is_export: bool = info['is_export']
# 查找原始函数
original_func: ir.Function | None = None
f: ir.Function
for f in module.functions:
if f.name == mangled_name:
original_func = f
break
if original_func is None:
continue
# 原函数改为 internal 链接(仅 wrapper 调用)
original_func.linkage = 'internal'
# 按装饰器列表从底到顶生成嵌套 wrapper
# Python 装饰器顺序:@a @b def f() => a(b(f))
# 执行顺序:先应用最靠近函数的 @b再应用 @a
current_func: ir.Function = original_func
all_wrapper_names: set[str] = set()
i: int
for i in range(len(decorators) - 1, -1, -1):
deco: dict = decorators[i]
is_outermost: bool = (i == 0)
deco_name: str = deco['name']
wrapper_name: str
if is_outermost:
wrapper_name = f"__decor_wrap_{mangled_name}"
else:
wrapper_name = f"__decor_wrap_{mangled_name}_{deco_name}"
all_wrapper_names.add(wrapper_name)
wrapper: ir.Function = _generate_single_wrapper(
module, current_func, wrapper_name,
func_name, deco, is_export=is_export,
is_outermost=is_outermost,
true_original_func=original_func if is_outermost else None
)
current_func = wrapper
# 更新 functions 映射,使原函数名指向最外层 wrapper
Gen.functions[mangled_name] = current_func
if func_name in Gen.functions:
Gen.functions[func_name] = current_func
# 替换所有非 wrapper 函数中对原函数的调用
_redirect_calls(module, original_func, current_func, all_wrapper_names)