290 lines
12 KiB
Python
290 lines
12 KiB
Python
import t, c
|
||
from stdint import *
|
||
import ast
|
||
import llvmlite
|
||
import memhub
|
||
import string
|
||
import viperlib
|
||
import lib.core.Handles.HandlesExpr as HandlesExpr
|
||
import lib.core.Handles.HandlesVar as HandlesVar
|
||
import lib.core.Handles.HandlesTranslator as HT
|
||
|
||
|
||
# ============================================================
|
||
# HandlesExprOps - 二元运算处理(模块级纯函数)
|
||
# ============================================================
|
||
|
||
# ============================================================
|
||
# 运算符重载支持
|
||
#
|
||
# 当二元/比较运算的 lhs 是结构体指针时,检查该类是否定义了
|
||
# 对应的 dunder 方法(如 __add__、__eq__),若有则生成方法调用
|
||
# 而非原生算术/比较指令。
|
||
# ============================================================
|
||
|
||
# BinOp 运算符 → dunder 方法名
|
||
def _binop_to_dunder(op: int) -> str:
|
||
"""将 BinOp 运算符映射到 dunder 方法名,无映射时返回 None"""
|
||
if op == ast.OpKind.Add: return "__add__"
|
||
if op == ast.OpKind.Sub: return "__sub__"
|
||
if op == ast.OpKind.Mult: return "__mul__"
|
||
if op == ast.OpKind.Div: return "__div__"
|
||
if op == ast.OpKind.Mod: return "__mod__"
|
||
if op == ast.OpKind.BitAnd: return "__and__"
|
||
if op == ast.OpKind.BitOr: return "__or__"
|
||
if op == ast.OpKind.BitXor: return "__xor__"
|
||
if op == ast.OpKind.LShift: return "__lshift__"
|
||
if op == ast.OpKind.RShift: return "__rshift__"
|
||
if op == ast.OpKind.FloorDiv: return "__floordiv__"
|
||
return None
|
||
|
||
# Compare 运算符 → dunder 方法名
|
||
def _cmpop_to_dunder(op: int) -> str:
|
||
"""将 Compare 运算符映射到 dunder 方法名,无映射时返回 None"""
|
||
if op == ast.OpKind.Eq: return "__eq__"
|
||
if op == ast.OpKind.Ne: return "__ne__"
|
||
if op == ast.OpKind.Lt: return "__lt__"
|
||
if op == ast.OpKind.Le: return "__le__"
|
||
if op == ast.OpKind.Gt: return "__gt__"
|
||
if op == ast.OpKind.Ge: return "__ge__"
|
||
return None
|
||
|
||
|
||
# ============================================================
|
||
# try_operator_overload - 尝试运算符重载
|
||
#
|
||
# 检查 lhs 是否为结构体指针,并在该类(含继承链)中查找对应的
|
||
# dunder 方法。找到则生成方法调用 lhs.__dunder__(rhs),返回
|
||
# 调用结果;未找到返回 None,调用方回退到原生运算。
|
||
#
|
||
# Args:
|
||
# pool: 内存池
|
||
# builder: IRBuilder
|
||
# mod: LLVM 模块
|
||
# lhs: 左操作数 Value(已求值)
|
||
# rhs: 右操作数 Value(已求值)
|
||
# op: OpKind 运算符
|
||
# trans: Translator 对象
|
||
# is_compare: 0=BinOp, 1=Compare(决定 dunder 映射表)
|
||
#
|
||
# Returns:
|
||
# 方法调用的 Value(成功),None(未重载)
|
||
# ============================================================
|
||
def try_operator_overload(pool: memhub.MemBuddy | t.CPtr,
|
||
builder: llvmlite.IRBuilder | t.CPtr,
|
||
mod: llvmlite.LLVMModule | t.CPtr,
|
||
lhs: llvmlite.Value | t.CPtr,
|
||
rhs: llvmlite.Value | t.CPtr,
|
||
op: int,
|
||
trans: HT.Translator | t.CPtr,
|
||
is_compare: int) -> llvmlite.Value | t.CPtr:
|
||
"""尝试运算符重载:lhs 是结构体指针时调用 dunder 方法,否则返回 None"""
|
||
if lhs is None or rhs is None:
|
||
return None
|
||
|
||
# 延迟导入避免循环依赖
|
||
import lib.core.Handles.HandlesStruct as HandlesStruct
|
||
import lib.core.Handles.HandlesExprCall as HandlesExprCall
|
||
|
||
# 映射运算符到 dunder 方法名
|
||
dunder: str = None
|
||
if is_compare != 0:
|
||
dunder = _cmpop_to_dunder(op)
|
||
else:
|
||
dunder = _binop_to_dunder(op)
|
||
if dunder is None:
|
||
return None
|
||
|
||
# 检查 lhs 是否为结构体指针
|
||
struct_ty: llvmlite.LLVMType | t.CPtr = HandlesStruct.get_struct_type_from_value(lhs)
|
||
if struct_ty is None:
|
||
return None
|
||
|
||
# 获取类名
|
||
class_name: str = HandlesStruct.get_class_name_by_type(pool, struct_ty)
|
||
if class_name is None:
|
||
return None
|
||
|
||
# 在继承链中查找 dunder 方法
|
||
lookup_name: t.CChar | t.CPtr = pool.alloc(128)
|
||
if lookup_name is None:
|
||
return None
|
||
viperlib.snprintf(lookup_name, 128, "%s.%s", class_name, dunder)
|
||
found_func: llvmlite.Function | t.CPtr = HandlesExprCall.find_func_in_module(mod, lookup_name)
|
||
|
||
search_class: str = class_name
|
||
if found_func is None:
|
||
cur_parent: str = HandlesStruct.get_parent_name(class_name)
|
||
while cur_parent is not None:
|
||
parent_lookup: t.CChar | t.CPtr = pool.alloc(128)
|
||
if parent_lookup is not None:
|
||
viperlib.snprintf(parent_lookup, 128, "%s.%s", cur_parent, dunder)
|
||
parent_func: llvmlite.Function | t.CPtr = HandlesExprCall.find_func_in_module(mod, parent_lookup)
|
||
if parent_func is not None:
|
||
found_func = parent_func
|
||
search_class = cur_parent
|
||
break
|
||
cur_parent = HandlesStruct.get_parent_name(cur_parent)
|
||
|
||
# 方法不在当前模块中时,检查类(含父类)是否有 SHA1(stub 可能未注入)
|
||
# 优先用类型指针定位 entry 获取 SHA1(规避跨模块同名 find_struct 找错)
|
||
cls_sha1: str = None
|
||
op_entry: HandlesStruct.StructEntry | t.CPtr = HandlesStruct.find_struct_by_type(struct_ty)
|
||
if op_entry is not None:
|
||
cls_sha1 = op_entry.ModuleSha1
|
||
if cls_sha1 is None:
|
||
cls_sha1 = HandlesStruct.get_struct_sha1(search_class)
|
||
if cls_sha1 is None:
|
||
cur_p: str = HandlesStruct.get_parent_name(search_class)
|
||
while cur_p is not None and cls_sha1 is None:
|
||
cls_sha1 = HandlesStruct.get_struct_sha1(cur_p)
|
||
cur_p = HandlesStruct.get_parent_name(cur_p)
|
||
|
||
# 既无函数定义也无 SHA1 → 该类没有此 dunder 方法,返回 None 回退原生运算
|
||
if found_func is None and cls_sha1 is None:
|
||
return None
|
||
|
||
# 构建 extra_args 数组(仅含 rhs 一个参数)
|
||
extra_args: t.CSizeT | t.CPtr = pool.alloc(8)
|
||
if extra_args is None:
|
||
return None
|
||
extra_args[0] = t.CSizeT(rhs)
|
||
|
||
# 调用方法: lhs.__dunder__(rhs)
|
||
return HandlesExprCall._call_method_on_ptr(
|
||
pool, builder, mod, search_class, dunder,
|
||
lhs, extra_args, 1, trans)
|
||
|
||
|
||
# ============================================================
|
||
# 翻译二元运算(自动类型提升)
|
||
# ============================================================
|
||
def translate_binop(pool: memhub.MemBuddy | t.CPtr,
|
||
builder: llvmlite.IRBuilder | t.CPtr,
|
||
mod: llvmlite.LLVMModule | t.CPtr,
|
||
node: ast.AST | t.CPtr,
|
||
trans: HT.Translator | t.CPtr = None) -> llvmlite.Value | t.CPtr:
|
||
"""翻译二元运算(自动类型提升)"""
|
||
binop: ast.BinOp | t.CPtr = (ast.BinOp | t.CPtr)(node)
|
||
if binop is None:
|
||
return None
|
||
|
||
lhs_node: ast.AST | t.CPtr = binop.left
|
||
rhs_node: ast.AST | t.CPtr = binop.right
|
||
op: int = binop.op
|
||
|
||
# 先翻译 rhs(只翻译一次,避免副作用重复)
|
||
rhs: llvmlite.Value | t.CPtr = HandlesExpr.translate_value(
|
||
builder, pool, mod, rhs_node, None, 0, trans)
|
||
if rhs is None:
|
||
return None
|
||
|
||
# === 运算符重载路径 1: lhs 是 Name 且对应结构体变量 ===
|
||
# 对于值类型变量(如 cnt: Counter),translate_value 会 load 返回 Struct 值,
|
||
# 但 dunder 方法需要 Ptr(Struct) 作为 self。直接用 alloca 指针尝试重载。
|
||
if lhs_node is not None and lhs_node.kind() == ast.ASTKind.Name and trans is not None:
|
||
nm: ast.Name | t.CPtr = (ast.Name | t.CPtr)(lhs_node)
|
||
if nm is not None and nm.id is not None:
|
||
lhs_alloca: llvmlite.Value | t.CPtr = HandlesVar.lookup_var(
|
||
trans.SymTab, nm.id)
|
||
if lhs_alloca is not None:
|
||
ovl_result: llvmlite.Value | t.CPtr = try_operator_overload(
|
||
pool, builder, mod, lhs_alloca, rhs, op, trans, 0)
|
||
if ovl_result is not None:
|
||
return ovl_result
|
||
|
||
# 正常翻译 lhs
|
||
lhs: llvmlite.Value | t.CPtr = HandlesExpr.translate_value(
|
||
builder, pool, mod, lhs_node, None, 0, trans)
|
||
if lhs is None:
|
||
return None
|
||
|
||
# === 运算符重载路径 2: lhs 是 Ptr(Struct)(如 Counter|t.CPtr 变量 load 后)===
|
||
ovl_result2: llvmlite.Value | t.CPtr = try_operator_overload(
|
||
pool, builder, mod, lhs, rhs, op, trans, 0)
|
||
if ovl_result2 is not None:
|
||
return ovl_result2
|
||
|
||
# 获取操作数类型信息
|
||
lhs_bits: int = HandlesExpr.get_llvm_type_bits(lhs.Ty)
|
||
rhs_bits: int = HandlesExpr.get_llvm_type_bits(rhs.Ty)
|
||
lhs_fbits: int = HandlesExpr.get_llvm_float_bits(lhs.Ty)
|
||
rhs_fbits: int = HandlesExpr.get_llvm_float_bits(rhs.Ty)
|
||
|
||
# 浮点运算: 任一操作数为浮点时,使用浮点指令(必须在指针算术之前检查,
|
||
# 因为 float 的 int bits 为 0,会被误认为指针)
|
||
if lhs_fbits != 0 or rhs_fbits != 0:
|
||
# 确定目标浮点类型(使用较大的位宽)
|
||
target_fbits: int = lhs_fbits
|
||
if rhs_fbits > target_fbits:
|
||
target_fbits = rhs_fbits
|
||
target_float_ty: llvmlite.LLVMType | t.CPtr = llvmlite.Double(pool)
|
||
if target_fbits == 32:
|
||
target_float_ty = llvmlite.Float(pool)
|
||
# 将两个操作数转换为目标浮点类型(int→float via si2fp, float→float via fpext/fptrunc)
|
||
lhs = HandlesExpr.coerce_to_type(builder, lhs, target_float_ty)
|
||
rhs = HandlesExpr.coerce_to_type(builder, rhs, target_float_ty)
|
||
if op == ast.OpKind.Add:
|
||
return llvmlite.build_fadd(builder, lhs, rhs)
|
||
if op == ast.OpKind.Sub:
|
||
return llvmlite.build_fsub(builder, lhs, rhs)
|
||
if op == ast.OpKind.Mult:
|
||
return llvmlite.build_fmul(builder, lhs, rhs)
|
||
if op == ast.OpKind.Div:
|
||
return llvmlite.build_fdiv(builder, lhs, rhs)
|
||
if op == ast.OpKind.Mod:
|
||
return llvmlite.build_frem(builder, lhs, rhs)
|
||
return None
|
||
|
||
# 指针算术: ptr + int 或 ptr - int → ptrtoint + add/sub + inttoptr
|
||
if (lhs_bits == 0 and rhs_bits != 0) or (lhs_bits != 0 and rhs_bits == 0):
|
||
if op == ast.OpKind.Add or op == ast.OpKind.Sub:
|
||
i64_ty: llvmlite.LLVMType | t.CPtr = llvmlite.Int64(pool)
|
||
if lhs_bits == 0:
|
||
# lhs 是指针,rhs 是整数
|
||
ptr_as_int: llvmlite.Value | t.CPtr = llvmlite.build_ptrtoint(builder, lhs, i64_ty)
|
||
int_val: llvmlite.Value | t.CPtr = HandlesExpr.coerce_to_type(builder, rhs, i64_ty)
|
||
if op == ast.OpKind.Add:
|
||
result: llvmlite.Value | t.CPtr = llvmlite.build_add(builder, ptr_as_int, int_val)
|
||
else:
|
||
result = llvmlite.build_sub(builder, ptr_as_int, int_val)
|
||
return llvmlite.build_inttoptr(builder, result, lhs.Ty)
|
||
else:
|
||
# rhs 是指针,lhs 是整数(仅 Add 支持交换)
|
||
ptr_as_int = llvmlite.build_ptrtoint(builder, rhs, i64_ty)
|
||
int_val = HandlesExpr.coerce_to_type(builder, lhs, i64_ty)
|
||
if op == ast.OpKind.Add:
|
||
result = llvmlite.build_add(builder, int_val, ptr_as_int)
|
||
else:
|
||
result = llvmlite.build_sub(builder, int_val, ptr_as_int)
|
||
return llvmlite.build_inttoptr(builder, result, rhs.Ty)
|
||
|
||
if lhs_bits > rhs_bits:
|
||
rhs = HandlesExpr.coerce_to_type(builder, rhs, lhs.Ty)
|
||
elif rhs_bits > lhs_bits:
|
||
lhs = HandlesExpr.coerce_to_type(builder, lhs, rhs.Ty)
|
||
|
||
if op == ast.OpKind.Add:
|
||
return llvmlite.build_add(builder, lhs, rhs)
|
||
elif op == ast.OpKind.Sub:
|
||
return llvmlite.build_sub(builder, lhs, rhs)
|
||
elif op == ast.OpKind.Mult:
|
||
return llvmlite.build_mul(builder, lhs, rhs)
|
||
elif op == ast.OpKind.Div:
|
||
return llvmlite.build_sdiv(builder, lhs, rhs)
|
||
elif op == ast.OpKind.FloorDiv:
|
||
return llvmlite.build_sdiv(builder, lhs, rhs)
|
||
elif op == ast.OpKind.Mod:
|
||
return llvmlite.build_srem(builder, lhs, rhs)
|
||
elif op == ast.OpKind.BitAnd:
|
||
return llvmlite.build_and(builder, lhs, rhs)
|
||
elif op == ast.OpKind.BitOr:
|
||
return llvmlite.build_or(builder, lhs, rhs)
|
||
elif op == ast.OpKind.BitXor:
|
||
return llvmlite.build_xor(builder, lhs, rhs)
|
||
elif op == ast.OpKind.LShift:
|
||
return llvmlite.build_shl(builder, lhs, rhs)
|
||
elif op == ast.OpKind.RShift:
|
||
return llvmlite.build_ashr(builder, lhs, rhs)
|
||
return None
|