Files
TransPyV/App/lib/core/Handles/HandlesFor.py
2026-07-19 13:18:46 +08:00

385 lines
16 KiB
Python
Raw 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.
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