231 lines
8.3 KiB
Python
231 lines
8.3 KiB
Python
import t, c
|
||
from stdint import *
|
||
import ast
|
||
import llvmlite
|
||
import memhub
|
||
import string
|
||
import stdio
|
||
import viperlib
|
||
import lib.core.Handles.HandlesTranslator as HT
|
||
import lib.core.Handles.HandlesExpr as HandlesExpr
|
||
import lib.core.Handles.HandlesExprCall as HandlesExprCall
|
||
import lib.core.Handles.HandlesFunctions as HandlesFunctions
|
||
import lib.core.Handles.HandlesClassDef as HandlesClassDef
|
||
|
||
|
||
# ============================================================
|
||
# HandlesBody - 语句分派(trans 单参模式,全方法调用)
|
||
#
|
||
# 所有语句类型通过 trans.XxxH.Handle(node) 分派到对应 Handle。
|
||
# 对应 TransPyC BodyHandle.HandleBodyLlvm 的 isinstance 分派。
|
||
# ============================================================
|
||
|
||
|
||
# ============================================================
|
||
# 语句翻译分派
|
||
# ============================================================
|
||
def translate_stmt(trans: HT.Translator | t.CPtr,
|
||
node: ast.AST | t.CPtr) -> int:
|
||
"""翻译单条语句,返回新增的变量数"""
|
||
if node is None:
|
||
return 0
|
||
k: int = node.kind()
|
||
|
||
if k == ast.ASTKind.Expr:
|
||
return translate_expr_stmt(trans, node)
|
||
elif k == ast.ASTKind.Assign:
|
||
return trans.AssignH.Handle(node)
|
||
elif k == ast.ASTKind.AnnAssign:
|
||
return trans.AnnAssignH.Handle(node)
|
||
elif k == ast.ASTKind.FunctionDef:
|
||
# 嵌套函数定义:提升为顶层函数 + 创建闭包
|
||
return HandlesFunctions.translate_nested_function_def(trans, node)
|
||
elif k == ast.ASTKind.Return:
|
||
return trans.ReturnH.Handle(node)
|
||
elif k == ast.ASTKind.If:
|
||
return trans.IfH.Handle(node)
|
||
elif k == ast.ASTKind.While:
|
||
return trans.WhileH.Handle(node)
|
||
elif k == ast.ASTKind.AugAssign:
|
||
return trans.AugAssignH.Handle(node)
|
||
elif k == ast.ASTKind.For:
|
||
return trans.ForH.Handle(node)
|
||
elif k == ast.ASTKind.ClassDef:
|
||
return HandlesClassDef.translate_class_def(trans, node)
|
||
elif k == ast.ASTKind.Import:
|
||
return trans.ImportsH.HandleImport(node)
|
||
elif k == ast.ASTKind.ImportFrom:
|
||
trans.ImportsH.HandleImportFromModule(node)
|
||
return trans.ImportsH.HandleImportFromNames(node)
|
||
elif k == ast.ASTKind.Global:
|
||
return translate_global(trans, node)
|
||
elif k == ast.ASTKind.Nonlocal:
|
||
return translate_nonlocal(trans, node)
|
||
elif k == ast.ASTKind.Pass:
|
||
return 0
|
||
elif k == ast.ASTKind.Break:
|
||
return translate_break(trans)
|
||
elif k == ast.ASTKind.Continue:
|
||
return translate_continue(trans)
|
||
return 0
|
||
|
||
|
||
# ============================================================
|
||
# 翻译 Global 语句: global x, y
|
||
#
|
||
# 将 names 中的变量名加入 _global_names 集合
|
||
# ============================================================
|
||
def translate_global(trans: HT.Translator | t.CPtr,
|
||
node: ast.AST | t.CPtr) -> int:
|
||
"""翻译 global 语句:记录 global 变量名"""
|
||
gn: ast.Global | t.CPtr = (ast.Global | t.CPtr)(node)
|
||
if gn is None:
|
||
return 0
|
||
names: list[ast.AST | t.CPtr] | t.CPtr = gn.names
|
||
if names is None:
|
||
return 0
|
||
n: t.CSizeT = names.__len__()
|
||
for i in range(n):
|
||
nm_node: ast.AST | t.CPtr = names.get(i)
|
||
if nm_node is not None and nm_node.kind() == ast.ASTKind.Name:
|
||
nm: ast.Name | t.CPtr = (ast.Name | t.CPtr)(nm_node)
|
||
if nm.id is not None:
|
||
HT.add_global_name(trans, nm.id)
|
||
return 0
|
||
|
||
|
||
# ============================================================
|
||
# 翻译 Nonlocal 语句: nonlocal x, y
|
||
#
|
||
# 将 names 中的变量名加入 _nonlocal_names 集合
|
||
# ============================================================
|
||
def translate_nonlocal(trans: HT.Translator | t.CPtr,
|
||
node: ast.AST | t.CPtr) -> int:
|
||
"""翻译 nonlocal 语句:记录 nonlocal 变量名"""
|
||
nl: ast.Nonlocal | t.CPtr = (ast.Nonlocal | t.CPtr)(node)
|
||
if nl is None:
|
||
return 0
|
||
names: list[ast.AST | t.CPtr] | t.CPtr = nl.names
|
||
if names is None:
|
||
return 0
|
||
n: t.CSizeT = names.__len__()
|
||
for i in range(n):
|
||
nm_node: ast.AST | t.CPtr = names.get(i)
|
||
if nm_node is not None and nm_node.kind() == ast.ASTKind.Name:
|
||
nm: ast.Name | t.CPtr = (ast.Name | t.CPtr)(nm_node)
|
||
if nm.id is not None:
|
||
HT.add_nonlocal_name(trans, nm.id)
|
||
return 0
|
||
|
||
|
||
# ============================================================
|
||
# 翻译表达式语句 Expr(value=Call(...))
|
||
# ============================================================
|
||
def translate_expr_stmt(trans: HT.Translator | t.CPtr,
|
||
node: ast.AST | t.CPtr) -> int:
|
||
"""翻译表达式语句(printf 调用等)"""
|
||
ex: ast.Expr | t.CPtr = (ast.Expr | t.CPtr)(node)
|
||
if ex is None:
|
||
return 0
|
||
call_node: ast.AST | t.CPtr = ex.value
|
||
if call_node is None:
|
||
return 0
|
||
ck: int = call_node.kind()
|
||
if ck != ast.ASTKind.Call:
|
||
return 0
|
||
|
||
cl: ast.Call | t.CPtr = (ast.Call | t.CPtr)(call_node)
|
||
func_node: ast.AST | t.CPtr = cl.func
|
||
if func_node is None:
|
||
return 0
|
||
|
||
func_name: str = HandlesExpr.get_func_name(func_node)
|
||
if func_name is None:
|
||
return 0
|
||
|
||
# printf / print 走特殊路径(print 映射到 printf)
|
||
if string.strcmp(func_name, "printf") == 0 or string.strcmp(func_name, "print") == 0:
|
||
trans.ExprCallH.HandlePrintfCall(cl)
|
||
else:
|
||
# 通用函数调用(返回值丢弃)
|
||
trans.ExprCallH.HandleCall(call_node)
|
||
return 0
|
||
|
||
|
||
# ============================================================
|
||
# 预扫描:为局部变量提前创建 alloca
|
||
# ============================================================
|
||
def pre_scan_allocas(trans: HT.Translator | t.CPtr,
|
||
node: ast.AST | t.CPtr) -> int:
|
||
"""预扫描语句中的 AnnAssign,提前创建 alloca
|
||
|
||
返回新增的变量数
|
||
"""
|
||
if node is None:
|
||
return 0
|
||
return trans.AnnAssignH.PreScan(node)
|
||
|
||
|
||
# ============================================================
|
||
# 翻译 break 语句
|
||
# ============================================================
|
||
def translate_break(trans: HT.Translator | t.CPtr) -> int:
|
||
"""翻译 break 语句:跳转到循环 end 块"""
|
||
builder: llvmlite.IRBuilder | t.CPtr = trans._cur_builder
|
||
func: llvmlite.Function | t.CPtr = trans._cur_func
|
||
pool: memhub.MemBuddy | t.CPtr = trans.Pool
|
||
|
||
if builder is None or func is None:
|
||
return 0
|
||
|
||
# 发射 br 到 break 目标
|
||
if trans._break_bb is not None:
|
||
if llvmlite.builder_cur_block_is_terminated(builder) == 0:
|
||
llvmlite.build_br(builder, trans._break_bb)
|
||
|
||
# 创建死代码 BB,用于后续语句(break 后面的代码不可达)
|
||
cnt: int = trans._label_counter
|
||
trans._label_counter = cnt + 1
|
||
name_buf: t.CChar | t.CPtr = pool.alloc(32)
|
||
viperlib.snprintf(name_buf, 32, "dead.%d", cnt)
|
||
dead_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
llvmlite.position_at_end(builder, dead_bb)
|
||
return 0
|
||
|
||
|
||
# ============================================================
|
||
# 翻译 continue 语句
|
||
# ============================================================
|
||
def translate_continue(trans: HT.Translator | t.CPtr) -> int:
|
||
"""翻译 continue 语句:跳转到循环 cond/incr 块"""
|
||
builder: llvmlite.IRBuilder | t.CPtr = trans._cur_builder
|
||
func: llvmlite.Function | t.CPtr = trans._cur_func
|
||
pool: memhub.MemBuddy | t.CPtr = trans.Pool
|
||
|
||
if builder is None or func is None:
|
||
return 0
|
||
|
||
# 发射 br 到 continue 目标
|
||
if trans._continue_bb is not None:
|
||
if llvmlite.builder_cur_block_is_terminated(builder) == 0:
|
||
llvmlite.build_br(builder, trans._continue_bb)
|
||
|
||
# 创建死代码 BB,用于后续语句
|
||
cnt: int = trans._label_counter
|
||
trans._label_counter = cnt + 1
|
||
name_buf: t.CChar | t.CPtr = pool.alloc(32)
|
||
viperlib.snprintf(name_buf, 32, "dead.%d", cnt)
|
||
dead_bb: llvmlite.BasicBlock | t.CPtr = llvmlite.create_block(pool, func, name_buf)
|
||
llvmlite.position_at_end(builder, dead_bb)
|
||
return 0
|
||
|
||
|
||
# ============================================================
|
||
# 获取语句类型名(调试用)
|
||
# ============================================================
|
||
def get_stmt_kind_name(node: ast.AST | t.CPtr) -> str:
|
||
"""获取语句类型名"""
|
||
if node is None:
|
||
return None
|
||
return node.type_name()
|