385 lines
16 KiB
Python
385 lines
16 KiB
Python
import t, c
|
||
from stdint import *
|
||
import ast
|
||
import llvmlite
|
||
import memhub
|
||
import string
|
||
import stdio
|
||
import viperlib
|
||
import lib.core.Handles.HandlesBase as HandlesBase
|
||
import lib.core.Handles.HandlesTranslator as HT
|
||
import lib.core.Handles.HandlesVar as HandlesVar
|
||
import lib.core.Handles.HandlesExpr as HandlesExpr
|
||
import lib.core.Handles.HandlesBody as HandlesBody
|
||
import lib.core.Handles.HandlesType as HandlesType
|
||
|
||
|
||
# ============================================================
|
||
# HandlesFor - for 循环语句处理(Mixin 继承模式)
|
||
#
|
||
# 支持 for i in range(start, stop, step) 模式:
|
||
# %i = alloca i32
|
||
# store i32 start, i32* %i
|
||
# br label %cond
|
||
# cond:
|
||
# %cur = load i32, i32* %i
|
||
# %cmp = icmp slt i32 %cur, stop
|
||
# br i1 %cmp, label %body, label %end
|
||
# body:
|
||
# ... body ...
|
||
# br label %incr
|
||
# incr:
|
||
# %cur2 = load i32, i32* %i
|
||
# %next = add i32 %cur2, step
|
||
# store i32 %next, i32* %i
|
||
# br label %cond
|
||
# end:
|
||
# ============================================================
|
||
|
||
|
||
@t.NoVTable
|
||
class ForHandle(HandlesBase.Mixin):
|
||
"""for 循环语句处理器:继承 Mixin 获得 Trans 回指针"""
|
||
|
||
def __init__(self, trans: HT.Translator | t.CPtr):
|
||
self.Trans = trans
|
||
|
||
# ============================================================
|
||
# Handle - 处理 for 语句,返回新增变量数
|
||
# ============================================================
|
||
def Handle(self, node: ast.AST | t.CPtr) -> int:
|
||
"""翻译 for i in range(...) 循环语句"""
|
||
if node is None:
|
||
return 0
|
||
|
||
trans: HT.Translator | t.CPtr = self.Trans
|
||
pool: memhub.MemBuddy | t.CPtr = trans.Pool
|
||
builder: llvmlite.IRBuilder | t.CPtr = trans._cur_builder
|
||
func: llvmlite.Function | t.CPtr = trans._cur_func
|
||
|
||
if builder is None or func is None:
|
||
return 0
|
||
|
||
for_node: ast.For | t.CPtr = (ast.For | t.CPtr)(node)
|
||
|
||
# 1. 获取循环变量名(仅支持 for i in range(...))
|
||
target: ast.AST | t.CPtr = for_node.target
|
||
if target is None:
|
||
return 0
|
||
if target.kind() != ast.ASTKind.Name:
|
||
return 0
|
||
target_nm: ast.Name | t.CPtr = (ast.Name | t.CPtr)(target)
|
||
var_name: str = target_nm.id
|
||
if var_name is None:
|
||
return 0
|
||
|
||
# 2. 解析迭代器:支持 range() 和指针迭代
|
||
iter_node: ast.AST | t.CPtr = for_node.iter
|
||
if iter_node is None:
|
||
return 0
|
||
|
||
# 检查是否是 range() 调用
|
||
is_range_iter: int = 0
|
||
if iter_node.kind() == ast.ASTKind.Call:
|
||
call_pre: ast.Call | t.CPtr = (ast.Call | t.CPtr)(iter_node)
|
||
fn_pre: str = HandlesExpr.get_func_name(call_pre.func)
|
||
if fn_pre is not None:
|
||
if string.strcmp(fn_pre, "range") == 0:
|
||
is_range_iter = 1
|
||
else:
|
||
err_msg: t.CChar | t.CPtr = pool.alloc(256)
|
||
if err_msg is not None:
|
||
viperlib.snprintf(err_msg, 256, "仅支持 range() 或指针迭代,got call '%s'", fn_pre)
|
||
HandlesType.fatal_error(iter_node, err_msg)
|
||
HandlesType.fatal_error(iter_node, "仅支持 range() 或指针迭代")
|
||
|
||
# 非范围迭代:走指针迭代路径
|
||
if is_range_iter == 0:
|
||
return self._handle_ptr_iter(for_node, var_name)
|
||
|
||
call: ast.Call | t.CPtr = (ast.Call | t.CPtr)(iter_node)
|
||
|
||
# 3. 解析 range 参数: range(stop) / range(start, stop) / range(start, stop, step)
|
||
args: list[ast.AST | t.CPtr] | t.CPtr = call.args
|
||
if args is None:
|
||
return 0
|
||
arg_count: t.CSizeT = args.__len__()
|
||
|
||
i32_ty: llvmlite.LLVMType | t.CPtr = llvmlite.Int32(pool)
|
||
start_val: llvmlite.Value | t.CPtr = llvmlite.const_int32(pool, 0)
|
||
stop_val: llvmlite.Value | t.CPtr = None
|
||
step_val: llvmlite.Value | t.CPtr = llvmlite.const_int32(pool, 1)
|
||
|
||
if arg_count == 1:
|
||
stop_val = HandlesExpr.translate_value(
|
||
builder, pool, trans.Module, args.get(0),
|
||
trans._funcs, trans._func_count, trans)
|
||
elif arg_count >= 2:
|
||
start_val = HandlesExpr.translate_value(
|
||
builder, pool, trans.Module, args.get(0),
|
||
trans._funcs, trans._func_count, trans)
|
||
stop_val = HandlesExpr.translate_value(
|
||
builder, pool, trans.Module, args.get(1),
|
||
trans._funcs, trans._func_count, trans)
|
||
if arg_count >= 3:
|
||
step_val = HandlesExpr.translate_value(
|
||
builder, pool, trans.Module, args.get(2),
|
||
trans._funcs, trans._func_count, trans)
|
||
|
||
if stop_val is None:
|
||
stop_val = llvmlite.const_int32(pool, 0)
|
||
if start_val is None:
|
||
start_val = llvmlite.const_int32(pool, 0)
|
||
if step_val is None:
|
||
step_val = llvmlite.const_int32(pool, 1)
|
||
|
||
# 4. 创建/查找循环变量 alloca
|
||
var_alloca: llvmlite.Value | t.CPtr = HandlesVar.lookup_var(
|
||
trans.SymTab, var_name)
|
||
new_vars: int = 0
|
||
if var_alloca is None:
|
||
var_alloca = llvmlite.build_alloca(builder, i32_ty)
|
||
if HandlesVar.define_var(
|
||
trans.SymTab, var_name, var_alloca) == 0:
|
||
new_vars = 1
|
||
|
||
# 5. 存储初始值 (类型对齐: start_val 可能是 i64,需截断到 i32)
|
||
init_val: llvmlite.Value | t.CPtr = start_val
|
||
if start_val is not None and start_val.Ty is not None:
|
||
start_bits: int = HandlesExpr.get_llvm_type_bits(start_val.Ty)
|
||
if start_bits != 0 and start_bits != 32:
|
||
init_val = llvmlite.build_trunc(builder, start_val, i32_ty)
|
||
llvmlite.build_store(builder, init_val, var_alloca)
|
||
|
||
# 6. 创建基本块: cond / body / incr / end(使用 trans._label_counter,不与 SSA 名共享)
|
||
cnt: int = trans._label_counter
|
||
trans._label_counter = cnt + 1
|
||
|
||
name_buf: t.CChar | t.CPtr = pool.alloc(32)
|
||
viperlib.snprintf(name_buf, 32, "for.cond.%d", cnt)
|
||
cond_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
viperlib.snprintf(name_buf, 32, "for.body.%d", cnt)
|
||
body_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
viperlib.snprintf(name_buf, 32, "for.incr.%d", cnt)
|
||
incr_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
viperlib.snprintf(name_buf, 32, "for.end.%d", cnt)
|
||
end_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
# 7. 跳转到 cond 块
|
||
llvmlite.build_br(builder, cond_bb)
|
||
|
||
# 8. cond 块: load i, icmp slt i, stop, cond_br body/end
|
||
llvmlite.position_at_end(builder, cond_bb)
|
||
cur_i: llvmlite.Value | t.CPtr = llvmlite.build_load(builder, i32_ty, var_alloca)
|
||
# 类型对齐: stop_val 可能是 i64 (如 range(strlen(s))),需将 cur_i 提升到 stop_val 类型
|
||
cmp_lhs: llvmlite.Value | t.CPtr = cur_i
|
||
cmp_rhs: llvmlite.Value | t.CPtr = stop_val
|
||
if stop_val is not None and stop_val.Ty is not None:
|
||
stop_bits: int = HandlesExpr.get_llvm_type_bits(stop_val.Ty)
|
||
cur_bits: int = HandlesExpr.get_llvm_type_bits(cur_i.Ty)
|
||
if stop_bits != 0 and cur_bits != 0 and stop_bits != cur_bits:
|
||
if cur_bits < stop_bits:
|
||
cmp_lhs = llvmlite.build_sext(builder, cur_i, stop_val.Ty)
|
||
else:
|
||
cmp_rhs = llvmlite.build_trunc(builder, stop_val, i32_ty)
|
||
cond_i1: llvmlite.Value | t.CPtr = llvmlite.build_icmp(
|
||
builder, llvmlite.ICMP_SLT, cmp_lhs, cmp_rhs)
|
||
llvmlite.build_cond_br(builder, cond_i1, body_bb, end_bb)
|
||
|
||
# 9. body 块: 翻译循环体,跳到 incr
|
||
llvmlite.position_at_end(builder, body_bb)
|
||
|
||
# 保存旧循环上下文,设置 break/continue 目标
|
||
old_break: llvmlite.BasicBlock | t.CPtr = trans._break_bb
|
||
old_continue: llvmlite.BasicBlock | t.CPtr = trans._continue_bb
|
||
trans._break_bb = end_bb
|
||
trans._continue_bb = incr_bb
|
||
|
||
body: list[ast.AST | t.CPtr] | t.CPtr = for_node.children
|
||
if body is not None:
|
||
body_count: t.CSizeT = body.__len__()
|
||
for bi in range(body_count):
|
||
stmt: ast.AST | t.CPtr = body.get(bi)
|
||
if stmt is not None:
|
||
HandlesBody.translate_stmt(trans, stmt)
|
||
|
||
# 恢复旧循环上下文
|
||
trans._break_bb = old_break
|
||
trans._continue_bb = old_continue
|
||
|
||
if llvmlite.builder_cur_block_is_terminated(builder) == 0:
|
||
llvmlite.build_br(builder, incr_bb)
|
||
|
||
# 10. incr 块: i = i + step, 跳回 cond
|
||
llvmlite.position_at_end(builder, incr_bb)
|
||
cur_i2: llvmlite.Value | t.CPtr = llvmlite.build_load(builder, i32_ty, var_alloca)
|
||
# 类型对齐: step_val 可能是 i64,需截断到 i32 与 cur_i2 类型一致
|
||
incr_step: llvmlite.Value | t.CPtr = step_val
|
||
if step_val is not None and step_val.Ty is not None:
|
||
step_bits: int = HandlesExpr.get_llvm_type_bits(step_val.Ty)
|
||
cur2_bits: int = HandlesExpr.get_llvm_type_bits(cur_i2.Ty)
|
||
if step_bits != 0 and cur2_bits != 0 and step_bits != cur2_bits:
|
||
if step_bits > cur2_bits:
|
||
incr_step = llvmlite.build_trunc(builder, step_val, i32_ty)
|
||
else:
|
||
incr_step = llvmlite.build_sext(builder, step_val, i32_ty)
|
||
next_i: llvmlite.Value | t.CPtr = llvmlite.build_add(builder, cur_i2, incr_step)
|
||
llvmlite.build_store(builder, next_i, var_alloca)
|
||
llvmlite.build_br(builder, cond_bb)
|
||
|
||
# 11. 定位到 end 块
|
||
llvmlite.position_at_end(builder, end_bb)
|
||
return new_vars
|
||
|
||
# ============================================================
|
||
# _handle_ptr_iter - 指针迭代: for x in ptr:
|
||
# 遍历指针,依赖隐式 index,直到解引用为空(null 终止符)
|
||
#
|
||
# 生成 IR 结构:
|
||
# %idx = alloca i32
|
||
# store i32 0, i32* %idx
|
||
# br label %cond
|
||
# cond:
|
||
# %i = load i32, i32* %idx
|
||
# %ep = getelementptr elem_ty, ptr_ty %ptr, i32 %i
|
||
# %ev = load elem_ty, elem_ty* %ep
|
||
# %null = icmp eq elem_ty %ev, 0
|
||
# br i1 %null, label %end, label %body
|
||
# body:
|
||
# store elem_ty %ev, elem_ty* %var
|
||
# ... 循环体 ...
|
||
# br label %incr
|
||
# incr:
|
||
# %next = add i32 %i, 1
|
||
# store i32 %next, i32* %idx
|
||
# br label %cond
|
||
# end:
|
||
# ============================================================
|
||
def _handle_ptr_iter(self, for_node: ast.For | t.CPtr,
|
||
var_name: str) -> int:
|
||
"""指针迭代: for x in ptr: 直到解引用为空"""
|
||
trans: HT.Translator | t.CPtr = self.Trans
|
||
pool: memhub.MemBuddy | t.CPtr = trans.Pool
|
||
builder: llvmlite.IRBuilder | t.CPtr = trans._cur_builder
|
||
func: llvmlite.Function | t.CPtr = trans._cur_func
|
||
|
||
# 翻译迭代器表达式,获取指针值
|
||
ptr_val: llvmlite.Value | t.CPtr = HandlesExpr.translate_value(
|
||
builder, pool, trans.Module, for_node.iter,
|
||
trans._funcs, trans._func_count, trans)
|
||
if ptr_val is None:
|
||
HandlesType.fatal_error(for_node, "指针迭代: 无法翻译迭代器表达式")
|
||
|
||
# 检查是否是指针类型
|
||
if HandlesExpr.is_ptr_type(ptr_val.Ty) == 0:
|
||
HandlesType.fatal_error(for_node, "指针迭代: 迭代器不是指针类型")
|
||
|
||
# 获取元素类型
|
||
elem_ty: llvmlite.LLVMType | t.CPtr = ptr_val.Ty.Pointee
|
||
if elem_ty is None:
|
||
HandlesType.fatal_error(for_node, "指针迭代: 无法获取元素类型")
|
||
|
||
# 元素类型必须是整数(用于 icmp eq 0 检查 null 终止符)
|
||
elem_bits: int = HandlesExpr.get_llvm_type_bits(elem_ty)
|
||
if elem_bits == 0:
|
||
HandlesType.fatal_error(for_node, "指针迭代: 元素类型不是整数,无法检查 null 终止符")
|
||
|
||
i32_ty: llvmlite.LLVMType | t.CPtr = llvmlite.Int32(pool)
|
||
|
||
# 创建循环变量 alloca(存储元素值)
|
||
var_alloca: llvmlite.Value | t.CPtr = HandlesVar.lookup_var(
|
||
trans.SymTab, var_name)
|
||
new_vars: int = 0
|
||
if var_alloca is None:
|
||
var_alloca = llvmlite.build_alloca(builder, elem_ty)
|
||
if HandlesVar.define_var(
|
||
trans.SymTab, var_name, var_alloca) == 0:
|
||
new_vars = 1
|
||
|
||
# 创建隐式 index 变量,初始为 0
|
||
idx_alloca: llvmlite.Value | t.CPtr = llvmlite.build_alloca(builder, i32_ty)
|
||
llvmlite.build_store(builder, llvmlite.const_int32(pool, 0), idx_alloca)
|
||
|
||
# 创建基本块: cond / body / incr / end
|
||
cnt: int = trans._label_counter
|
||
trans._label_counter = cnt + 1
|
||
|
||
name_buf: t.CChar | t.CPtr = pool.alloc(32)
|
||
viperlib.snprintf(name_buf, 32, "ptr.cond.%d", cnt)
|
||
cond_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
viperlib.snprintf(name_buf, 32, "ptr.body.%d", cnt)
|
||
body_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
viperlib.snprintf(name_buf, 32, "ptr.incr.%d", cnt)
|
||
incr_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
viperlib.snprintf(name_buf, 32, "ptr.end.%d", cnt)
|
||
end_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
|
||
# 跳转到 cond
|
||
llvmlite.build_br(builder, cond_bb)
|
||
|
||
# cond 块: load index, GEP, load elem, icmp eq 0
|
||
llvmlite.position_at_end(builder, cond_bb)
|
||
cur_idx: llvmlite.Value | t.CPtr = llvmlite.build_load(builder, i32_ty, idx_alloca)
|
||
elem_ptr: llvmlite.Value | t.CPtr = llvmlite.build_gep(
|
||
builder, elem_ty, ptr_val, cur_idx)
|
||
cur_elem: llvmlite.Value | t.CPtr = llvmlite.build_load(builder, elem_ty, elem_ptr)
|
||
null_val: llvmlite.Value | t.CPtr = llvmlite.const_int(pool, elem_bits, 0)
|
||
is_null: llvmlite.Value | t.CPtr = llvmlite.build_icmp(
|
||
builder, llvmlite.ICMP_EQ, cur_elem, null_val)
|
||
llvmlite.build_cond_br(builder, is_null, end_bb, body_bb)
|
||
|
||
# body 块: store elem to var, 翻译循环体
|
||
llvmlite.position_at_end(builder, body_bb)
|
||
llvmlite.build_store(builder, cur_elem, var_alloca)
|
||
|
||
# 保存旧循环上下文,设置 break/continue 目标
|
||
old_break: llvmlite.BasicBlock | t.CPtr = trans._break_bb
|
||
old_continue: llvmlite.BasicBlock | t.CPtr = trans._continue_bb
|
||
trans._break_bb = end_bb
|
||
trans._continue_bb = incr_bb
|
||
|
||
body: list[ast.AST | t.CPtr] | t.CPtr = for_node.children
|
||
if body is not None:
|
||
body_count: t.CSizeT = body.__len__()
|
||
for bi in range(body_count):
|
||
stmt: ast.AST | t.CPtr = body.get(bi)
|
||
if stmt is not None:
|
||
HandlesBody.translate_stmt(trans, stmt)
|
||
|
||
# 恢复旧循环上下文
|
||
trans._break_bb = old_break
|
||
trans._continue_bb = old_continue
|
||
|
||
if llvmlite.builder_cur_block_is_terminated(builder) == 0:
|
||
llvmlite.build_br(builder, incr_bb)
|
||
|
||
# incr 块: index++, br cond
|
||
llvmlite.position_at_end(builder, incr_bb)
|
||
one_val: llvmlite.Value | t.CPtr = llvmlite.const_int32(pool, 1)
|
||
next_idx: llvmlite.Value | t.CPtr = llvmlite.build_add(builder, cur_idx, one_val)
|
||
llvmlite.build_store(builder, next_idx, idx_alloca)
|
||
llvmlite.build_br(builder, cond_bb)
|
||
|
||
# 定位到 end 块
|
||
llvmlite.position_at_end(builder, end_bb)
|
||
return new_vars
|
||
|
||
|
||
# ============================================================
|
||
# NewForHandle - 工厂函数
|
||
# ============================================================
|
||
def NewForHandle(pool: memhub.MemBuddy | t.CPtr,
|
||
trans: HT.Translator | t.CPtr) -> ForHandle | t.CPtr:
|
||
h: ForHandle | t.CPtr = pool.alloc(ForHandle.__sizeof__())
|
||
if h is None:
|
||
return None
|
||
string.memset(h, 0, ForHandle.__sizeof__())
|
||
h.Trans = trans
|
||
return h
|