from __future__ import annotations from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from lib.core.translator import Translator from lib.core.LlvmCodeGenerator import LlvmCodeGenerator import ast import re import llvmlite.ir as ir from lib.core.Handles.HandlesBase import BaseHandle from lib.constants.config import mode as _config_mode from lib.includes.t import CTypeRegistry as _CR class ExprHandle(BaseHandle): # list 元素类型 → (字节大小, LLVM 元素类型) 映射表 # 用于 _HandleListLlvm 中 alloc_size / byte_offset / elem_ptr 类型统一查表, # 消除原代码三处重复的 if-elif 链。 _ELEM_TYPE_INFO: dict[str, tuple[int, ir.Type]] = { 'i8': (1, ir.IntType(8)), 'i16': (2, ir.IntType(16)), 'i32': (4, ir.IntType(32)), 'float': (4, ir.FloatType()), 'double': (8, ir.DoubleType()), } def __init__(self, translator: "Translator") -> None: super().__init__(translator) def GetOpSymbol(self, Op: ast.cmpop | ast.operator | ast.unaryoperator) -> str | None: return self.Trans.ExprUtils.GetOpSymbol(Op) def GetUnaryOpSymbol(self, Op: ast.unaryoperator) -> str | None: return self.Trans.ExprUtils.GetUnaryOpSymbol(Op) def GetComparatorSymbol(self, Op: ast.cmpop) -> str | None: return self.Trans.ExprUtils.GetComparatorSymbol(Op) def _get_var_class(self, node: ast.AST, Gen: LlvmCodeGenerator) -> str | None: return self.Trans.ExprUtils._get_var_class(node, Gen) def _try_operator_overLoad(self, ClassName: str, op_name: str, obj_val: ir.Value, Gen: LlvmCodeGenerator, other_val: ir.Value | None = None) -> ir.Value | None: return self.Trans.ExprUtils._try_operator_overLoad(ClassName, op_name, obj_val, Gen, other_val) def _get_llvm_member_offset(self, field_name: str, ClassName: str, Gen: LlvmCodeGenerator) -> int | None: return self.Trans.ExprAttrHandle._get_llvm_member_offset(field_name, ClassName, Gen) def HandleExprLlvm(self, Node: ast.AST, VarType: ir.Type | str | None = None) -> ir.Value | None: Gen: LlvmCodeGenerator = self.Trans.LlvmGen if not Gen or not Gen.builder: return None if isinstance(Node, ast.Constant): return self._HandleConstantLlvm(Node, VarType) elif isinstance(Node, ast.Name): return self._HandleNameLlvm(Node, VarType) elif isinstance(Node, ast.BinOp): return self.Trans.ExprOpsHandle._HandleBinOpLlvm(Node, VarType) elif isinstance(Node, ast.BoolOp): return self.Trans.ExprOpsHandle._HandleBoolOpLlvm(Node) elif isinstance(Node, ast.UnaryOp): return self.Trans.ExprOpsHandle._HandleUnaryOpLlvm(Node) elif isinstance(Node, ast.Call): return self.Trans.ExprCallHandle._HandleCallLlvm(Node) elif isinstance(Node, ast.Compare): return self.Trans.ExprOpsHandle._HandleCompareLlvm(Node) elif isinstance(Node, ast.Attribute): return self.Trans.ExprAttrHandle._HandleAttributeLlvm(Node) elif isinstance(Node, ast.Subscript): return self.Trans.ExprAttrHandle._HandleSubscriptLlvm(Node) elif isinstance(Node, ast.IfExp): return self.Trans.ExprLambdaHandle._HandleIfExpLlvm(Node) elif isinstance(Node, ast.NamedExpr): return self.Trans.ExprLambdaHandle._HandleNamedExprLlvm(Node) elif isinstance(Node, ast.JoinedStr): return self.Trans.ExprFormatHandle._HandleJoinedStrLlvm(Node) elif isinstance(Node, ast.Lambda): return self.Trans.ExprLambdaHandle._HandleLambdaLlvm(Node) elif isinstance(Node, ast.List): return self._HandleListLlvm(Node, VarType) return ir.Constant(ir.IntType(32), 0) def _HandleListLlvm(self, Node: ast.List, VarType: ir.Type | str | None = None) -> ir.Value | None: Gen: LlvmCodeGenerator = self.Trans.LlvmGen elements: list[ast.expr] = Node.elts if not elements: return ir.Constant(ir.IntType(8).as_pointer(), None) array_len: int = len(elements) elem_type: str = 'i8' if VarType and isinstance(VarType, str) and 'list[' in VarType: match: re.Match | None = re.search(r'list\[([^,\]]+)(?:,\s*(\d+))?\]', VarType) if match: type_str: str = match.group(1).strip() if match.group(2): array_len = int(match.group(2)) if 'CChar' in type_str or 'char' in type_str.lower(): elem_type = 'i8' elif 'CFloat' in type_str or type_str.lower() == 'float': elem_type = 'float' elif 'CDouble' in type_str or type_str.lower() == 'double': elem_type = 'double' elif 'CInt' in type_str or ('int' in type_str.lower() and 'unsigned' not in type_str.lower()): elem_type = 'i32' elif 'CUnsignedInt' in type_str or 'unsigned' in type_str.lower(): elem_type = 'i32' elif 'CUnsignedChar' in type_str or 'unsigned char' in type_str.lower(): elem_type = 'i8' elif 'CShort' in type_str or 'short' in type_str.lower(): elem_type = 'i16' elif 'CUnsignedShort' in type_str: elem_type = 'i16' elif elements and all(isinstance(e, (ast.Constant,)) and isinstance(getattr(e, 'value', None), float) for e in elements): elem_type = 'float' alloc_size: int = array_len elem_info: tuple[int, ir.Type] = self._ELEM_TYPE_INFO.get(elem_type, (1, ir.IntType(8))) elem_size: int = elem_info[0] elem_llvm_type: ir.Type = elem_info[1] alloc_size *= elem_size malloc_fn: Any = None malloc_arg_type: ir.IntType = ir.IntType(64) for fn in Gen.module.global_values: if fn.name == 'malloc': malloc_fn = fn func_ftype: Any = getattr(fn, 'ftype', None) if func_ftype and getattr(func_ftype, 'args', None): if func_ftype.args: malloc_arg_type = fn.ftype.args[0] break if not malloc_fn: try: malloc_fn_type: ir.FunctionType = ir.FunctionType(ir.PointerType(ir.IntType(8)), [ir.IntType(64)]) malloc_fn = ir.Function(Gen.module, malloc_fn_type, name='malloc') Gen.functions['malloc'] = malloc_fn except Exception: # 回退:malloc 函数创建失败时返回 null return ir.Constant(ir.IntType(8).as_pointer(), None) alloc_ptr: Any = Gen.builder.call(malloc_fn, [ir.Constant(malloc_arg_type, alloc_size)]) cast_ptr: Any = Gen.builder.bitcast(alloc_ptr, ir.PointerType(ir.IntType(8))) for i, elem in enumerate(elements): if i >= array_len: break elem_val: Any = self.HandleExprLlvm(elem, VarType=elem_type) if elem_val is None: continue byte_offset: int = i if elem_type == 'i32': byte_offset = i * 4 elif elem_type == 'i16': byte_offset = i * 2 elif elem_type == 'float': byte_offset = i * 4 elif elem_type == 'double': byte_offset = i * 8 ptr_as_int: Any = Gen.builder.ptrtoint(cast_ptr, ir.IntType(64)) new_ptr_val: Any = Gen.builder.add(ptr_as_int, ir.Constant(ir.IntType(64), byte_offset)) if elem_type == 'i32': elem_ptr: Any = Gen.builder.inttoptr(new_ptr_val, ir.PointerType(ir.IntType(32))) elif elem_type == 'i16': elem_ptr = Gen.builder.inttoptr(new_ptr_val, ir.PointerType(ir.IntType(16))) elif elem_type == 'float': elem_ptr = Gen.builder.inttoptr(new_ptr_val, ir.PointerType(ir.FloatType())) elif elem_type == 'double': elem_ptr = Gen.builder.inttoptr(new_ptr_val, ir.PointerType(ir.DoubleType())) else: elem_ptr = Gen.builder.inttoptr(new_ptr_val, ir.PointerType(ir.IntType(8))) if elem_type == 'i8': if isinstance(elem_val.type, ir.IntType) and elem_val.type.width > 8: elem_val = Gen.builder.trunc(elem_val, ir.IntType(8), name=f"trunc_elem_{i}") elif isinstance(elem_val.type, ir.PointerType): elem_val = Gen.builder.ptrtoint(elem_val, ir.IntType(64), name=f"ptrtoint_elem_{i}") elem_val = Gen.builder.trunc(elem_val, ir.IntType(8), name=f"trunc_elem_{i}") elif elem_type == 'i32': if isinstance(elem_val.type, ir.IntType): if elem_val.type.width < 32: elem_val = Gen.builder.zext(elem_val, ir.IntType(32), name=f"zext_elem_{i}") elif elem_val.type.width > 32: elem_val = Gen.builder.trunc(elem_val, ir.IntType(32), name=f"trunc_elem_{i}") elif isinstance(elem_val.type, ir.PointerType): elem_val = Gen.builder.ptrtoint(elem_val, ir.IntType(32), name=f"ptrtoint_elem_{i}") elif elem_type == 'i16': if isinstance(elem_val.type, ir.IntType): if elem_val.type.width < 16: elem_val = Gen.builder.zext(elem_val, ir.IntType(16), name=f"zext_elem_{i}") elif elem_val.type.width > 16: elem_val = Gen.builder.trunc(elem_val, ir.IntType(16), name=f"trunc_elem_{i}") elif elem_type == 'float': if isinstance(elem_val.type, ir.DoubleType): elem_val = Gen.builder.fptrunc(elem_val, ir.FloatType(), name=f"fptrunc_elem_{i}") elif isinstance(elem_val.type, ir.IntType): elem_val = Gen.builder.sitofp(elem_val, ir.FloatType(), name=f"sitofp_elem_{i}") elif elem_type == 'double': if isinstance(elem_val.type, ir.FloatType): elem_val = Gen.builder.fpext(elem_val, ir.DoubleType(), name=f"fpext_elem_{i}") elif isinstance(elem_val.type, ir.IntType): elem_val = Gen.builder.sitofp(elem_val, ir.DoubleType(), name=f"sitofp_elem_{i}") Gen.builder.store(elem_val, elem_ptr) if elem_type in ('float', 'double'): return Gen.builder.bitcast(alloc_ptr, ir.PointerType(elem_llvm_type), name=f"list_{elem_type}_ptr") return cast_ptr def _HandleConstantLlvm(self, Node: ast.Constant, VarType: ir.Type | str | None = None) -> ir.Value: Gen: LlvmCodeGenerator = self.Trans.LlvmGen Value: Any = Node.value if Value is None: if VarType is not None and isinstance(VarType, ir.PointerType): return ir.Constant(VarType, None) return ir.Constant(ir.IntType(8).as_pointer(), None) elif isinstance(Value, bool): return ir.Constant(ir.IntType(32), 1 if Value else 0) elif isinstance(Value, int): if VarType is not None: if isinstance(VarType, ir.IntType): return ir.Constant(VarType, Value & ((1 << VarType.width) - 1)) if isinstance(VarType, ir.PointerType): return ir.Constant(ir.IntType(64), Value) if isinstance(VarType, (ir.FloatType, ir.DoubleType)): return ir.Constant(VarType, float(Value)) if isinstance(VarType, str) and '*' in VarType: return ir.Constant(ir.IntType(64), Value) if isinstance(VarType, str): resolved: Any = _CR.ResolveName(VarType) or _CR.LLVMToCType(VarType) if resolved: ctype_cls: type = resolved[0] if isinstance(resolved, tuple) else resolved inst: Any = ctype_cls() size: int = getattr(inst, 'Size', 32) is_unsigned: bool = getattr(inst, 'IsSigned', True) is False mask: int = (1 << size) - 1 return ir.Constant(ir.IntType(size), Value & mask) llvm_match: re.Match | None = re.match(r'^i(\d+)$', VarType) if llvm_match: width: int = int(llvm_match.group(1)) return ir.Constant(ir.IntType(width), Value & ((1 << width) - 1)) if Value < 0 or Value > 0xFFFFFFFF: return ir.Constant(ir.IntType(64), Value) return ir.Constant(ir.IntType(32), Value & 0xFFFFFFFF) elif isinstance(Value, float): if VarType and isinstance(VarType, ir.FloatType): return ir.Constant(VarType, Value) elif VarType and isinstance(VarType, ir.DoubleType): return ir.Constant(VarType, Value) return ir.Constant(ir.DoubleType(), Value) elif isinstance(Value, str): is_wide: bool = getattr(Node, 'kind', None) == 'u' # 引号类型推断:单引号单字符 → i8,双引号/多字符/三重 → i8*,空 → 0(None) quote_char: str | None = None is_triple: bool = False if not is_wide: col_offset: int = getattr(Node, 'col_offset', 0) node_lineno: int = getattr(Node, 'lineno', 0) # AST 的 lineno 可能不准确(常量节点可能指向下一行), # 因此在 lineno 和 lineno-1 两行中查找引号 for try_lineno in (node_lineno, node_lineno - 1): source_line: str | None = Gen._get_source_line(try_lineno) if source_line and col_offset < len(source_line): ch: str = source_line[col_offset] if ch in ("'", '"'): quote_char = ch if col_offset + 2 < len(source_line) and source_line[col_offset:col_offset + 3] in ("'''", '"""'): is_triple = True break # 空字符串 → 创建指向 \0 的全局变量 (而非 null, 避免 strcmp 解引用空指针崩溃) # 走与非空字符串相同的路径, 由下方 else 分支处理 # 单引号单字符(非三重)→ i8 is_char_literal: bool = (not is_triple) and (quote_char == "'") and (len(Value) == 1) and not is_wide if is_char_literal: if VarType is not None and isinstance(VarType, ir.PointerType): pass # 类型不匹配:目标是 ptr 但值是 char,走 string 路径 elif isinstance(VarType, ir.IntType): return ir.Constant(VarType, ord(Value[0])) elif isinstance(VarType, str): resolved = _CR.ResolveName(VarType) or _CR.LLVMToCType(VarType) if resolved: ctype_cls = resolved[0] if isinstance(resolved, tuple) else resolved inst = ctype_cls() size = getattr(inst, 'Size', 32) return ir.Constant(ir.IntType(size), ord(Value[0])) llvm_match = re.match(r'^i(\d+)$', VarType) if llvm_match: return ir.Constant(ir.IntType(int(llvm_match.group(1))), ord(Value[0])) elif VarType is not None: return ir.Constant(ir.IntType(8), ord(Value[0])) else: # 无类型提示:单引号单字符默认为 i8 return ir.Constant(ir.IntType(8), ord(Value[0])) if is_wide: str_value: str = Value + '\x00' wide_chars: list[int] = [ord(c) for c in str_value] target_count: int = len(wide_chars) str_type: ir.ArrayType = ir.ArrayType(ir.IntType(16), target_count) gv_name: str = f"str_const_{Gen.string_const_counter}" gv: ir.GlobalVariable = ir.GlobalVariable(Gen.module, str_type, name=gv_name) gv.initializer = ir.Constant(str_type, wide_chars) gv.linkage = 'internal' Gen.string_const_counter += 1 ptr: Any = Gen.builder.gep(gv, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), 0)], name=f"{gv_name}_cast") return ptr else: str_value = Value + '\x00' str_bytes: bytes = str_value.encode('utf-8') target_count = len(str_bytes) if isinstance(VarType, ir.ArrayType) and isinstance(VarType.element, ir.IntType) and VarType.element.width == 8: target_count = VarType.count if len(str_bytes) < target_count: str_bytes = str_bytes + b'\x00' * (target_count - len(str_bytes)) else: str_bytes = str_bytes[:target_count] str_type = ir.ArrayType(ir.IntType(8), target_count) gv_name = f"str_const_{Gen.string_const_counter}" gv = ir.GlobalVariable(Gen.module, str_type, name=gv_name) gv.initializer = ir.Constant(str_type, bytearray(str_bytes)) gv.linkage = 'internal' Gen.string_const_counter += 1 ptr = Gen.builder.gep(gv, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), 0)], name=f"{gv_name}_cast") return ptr return ir.Constant(ir.IntType(32), 0) def _HandleNameLlvm(self, Node: ast.Name, VarType: ir.Type | str | None = None) -> ir.Value: Gen: LlvmCodeGenerator = self.Trans.LlvmGen VarName: str = Node.id define_constants: dict[str, Any] = getattr(Gen, '_define_constants', {}) if VarName in define_constants: val: Any = define_constants[VarName] if isinstance(val, int): # Use VarType hint if available to avoid unnecessary type mismatch if VarType is not None and isinstance(VarType, ir.IntType): return ir.Constant(VarType, val & ((1 << VarType.width) - 1)) if VarName in Gen.var_signedness and Gen.var_signedness[VarName]: if val < -(1 << 31) or val > 0xFFFFFFFF: return ir.Constant(ir.IntType(64), val & 0xFFFFFFFFFFFFFFFF) return ir.Constant(ir.IntType(32), val & 0xFFFFFFFF) if val < -(1 << 31) or val > 0xFFFFFFFF: return ir.Constant(ir.IntType(64), val) return ir.Constant(ir.IntType(32), val & 0xFFFFFFFF) elif isinstance(val, float): return ir.Constant(ir.DoubleType(), val) elif isinstance(val, str): return Gen._create_string_global(val) # 如果 _define_constants 中没有,尝试从符号表查找 if VarName not in getattr(Gen, '_define_constants', {}): try: sym_info: Any = self.translator.SymbolTable.lookup(VarName) if sym_info and sym_info.IsDefine and sym_info.DefineValue is not None: val = sym_info.DefineValue if isinstance(val, int): # Use VarType hint if available if VarType is not None and isinstance(VarType, ir.IntType): return ir.Constant(VarType, val & ((1 << VarType.width) - 1)) if val < -(1 << 31) or val > 0xFFFFFFFF: return ir.Constant(ir.IntType(64), val) return ir.Constant(ir.IntType(32), val & 0xFFFFFFFF) elif isinstance(val, float): return ir.Constant(ir.DoubleType(), val) elif isinstance(val, str): return Gen._create_string_global(val) except Exception as _e: if _config_mode == "strict": self.Trans.LogWarning(f"异常被忽略: {_e}") # 优先检查 variables(alloca),因为变量可能已被 += 等操作迁移到 variables if VarName in Gen.variables and Gen.variables[VarName] is not None: VarPtr: Any = Gen.variables[VarName] if isinstance(VarPtr, ir.GlobalVariable) and isinstance(VarPtr.type, ir.PointerType) and isinstance(VarPtr.type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)): if VarName in Gen.global_struct_class: Gen.var_struct_class[VarName] = Gen.global_struct_class[VarName] return VarPtr if isinstance(VarPtr.type, ir.PointerType) and isinstance(VarPtr.type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)): if VarName in Gen.global_struct_class: Gen.var_struct_class[VarName] = Gen.global_struct_class[VarName] return VarPtr loaded: ir.Value = Gen._load(VarPtr, name=VarName) # 大端局部变量:读取时 bswap 还原 loaded = Gen._apply_bswap_if_big(loaded, getattr(Gen, 'local_var_byteorders', {}).get(VarName, ""), f"bswap_load_{VarName}") return loaded if VarName in Gen._reg_values: return Gen._reg_values[VarName] if VarName in Gen._direct_values: return Gen._direct_values[VarName] if VarName in Gen.global_vars and VarName in Gen.module.globals: GVar: Any = Gen.module.globals[VarName] Gen.variables[VarName] = GVar if isinstance(GVar, ir.Function): return GVar if isinstance(GVar.type, ir.PointerType) and isinstance(GVar.type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)): return GVar return Gen._load(GVar, name=VarName) if VarName in Gen.module.globals: GVar = Gen.module.globals[VarName] Gen.variables[VarName] = GVar if isinstance(GVar, ir.Function): return GVar if VarName in Gen.global_struct_class: Gen.var_struct_class[VarName] = Gen.global_struct_class[VarName] if isinstance(GVar.type, ir.PointerType) and isinstance(GVar.type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)): return GVar return Gen._load(GVar, name=VarName) if VarName == 'self': return Gen.GetVarPtr('self') if Gen._has_function(VarName): func: Any = Gen._get_function(VarName) if func: return Gen.builder.bitcast(func, ir.IntType(8).as_pointer(), name=f"funcptr_{VarName}") if VarName in Gen.class_methods: return ir.Constant(ir.IntType(8).as_pointer(), None) if VarName == 'True' or VarName == 'False': return ir.Constant(ir.IntType(32), 1 if VarName == 'True' else 0) SymInfo: Any = self.Trans.SymbolTable.lookup(VarName) if SymInfo: if SymInfo.IsEnumMember and isinstance(SymInfo.value, int): return ir.Constant(ir.IntType(32), SymInfo.value) if SymInfo.IsFunction or SymInfo.IsFuncPtr: MangledName: str = Gen._mangle_func_name(VarName) if MangledName in Gen.module.globals: g: Any = Gen.module.globals[MangledName] if isinstance(g, ir.Function): Gen.functions[VarName] = g return Gen.builder.bitcast(g, ir.IntType(8).as_pointer(), name=f"funcptr_{VarName}") for sha1 in Gen.ModuleSha1Map.values(): prefixed: str = f"{sha1}.{VarName}" if prefixed in Gen.module.globals: g = Gen.module.globals[prefixed] if isinstance(g, ir.Function): Gen.functions[VarName] = g return Gen.builder.bitcast(g, ir.IntType(8).as_pointer(), name=f"funcptr_{VarName}") t_c_imported: dict[str, tuple[str, str]] = getattr(self.Trans, '_t_c_imported_names', {}) if VarName in t_c_imported: src_module: str src_name: str src_module, src_name = t_c_imported[VarName] virtual_attr: ast.Attribute = ast.Attribute( value=ast.Name(id=src_module, ctx=ast.Load()), attr=src_name, ctx=ast.Load() ) ast.copy_location(virtual_attr, Node) return self.Trans.ExprAttrHandle._HandleAttributeLlvm(virtual_attr) return ir.Constant(ir.IntType(32), 0)