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