Files
TransPyC/lib/core/Handles/HandlesExpr.py

430 lines
23 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.
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from lib.core.translator import Translator
from lib.core.Handles.HandlesBase import BaseHandle
import ast
import llvmlite.ir as ir
class ExprHandle(BaseHandle):
def __init__(self, translator: "Translator"):
super().__init__(translator)
def GetOpSymbol(self, Op):
return self.Trans.ExprUtils.GetOpSymbol(Op)
def GetUnaryOpSymbol(self, Op):
return self.Trans.ExprUtils.GetUnaryOpSymbol(Op)
def GetComparatorSymbol(self, Op):
return self.Trans.ExprUtils.GetComparatorSymbol(Op)
def _get_var_class(self, node, Gen):
return self.Trans.ExprUtils._get_var_class(node, Gen)
def _try_operator_overLoad(self, ClassName, op_name, obj_val, Gen, other_val=None):
return self.Trans.ExprUtils._try_operator_overLoad(ClassName, op_name, obj_val, Gen, other_val)
def _get_llvm_member_offset(self, field_name, ClassName, Gen):
return self.Trans.ExprAttrHandle._get_llvm_member_offset(field_name, ClassName, Gen)
def HandleExprLlvm(self, Node, VarType=None):
Gen = 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):
in_str_method = getattr(Gen, 'func', None) and getattr(Gen.func, 'name', '').endswith('.__str__')
return self.Trans.ExprFormatHandle._HandleJoinedStrLlvm(Node, return_str=in_str_method)
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, VarType=None):
Gen = self.Trans.LlvmGen
elements = Node.elts
if not elements:
return ir.Constant(ir.IntType(8).as_pointer(), None)
array_len = len(elements)
elem_type = 'i8'
if VarType and isinstance(VarType, str) and 'list[' in VarType:
import re
match = re.search(r'list\[([^,\]]+)(?:,\s*(\d+))?\]', VarType)
if match:
type_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 = array_len
if elem_type == 'i32':
alloc_size *= 4
elif elem_type == 'i16':
alloc_size *= 2
elif elem_type == 'float':
alloc_size *= 4
elif elem_type == 'double':
alloc_size *= 8
malloc_fn = None
malloc_arg_type = ir.IntType(64)
for fn in Gen.module.global_values:
if fn.name == 'malloc':
malloc_fn = fn
func_ftype = 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.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 = Gen.builder.call(malloc_fn, [ir.Constant(malloc_arg_type, alloc_size)])
cast_ptr = Gen.builder.bitcast(alloc_ptr, ir.PointerType(ir.IntType(8)))
for i, elem in enumerate(elements):
if i >= array_len:
break
elem_val = self.HandleExprLlvm(elem, VarType=elem_type)
if elem_val is None:
continue
byte_offset = 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 = Gen.builder.ptrtoint(cast_ptr, ir.IntType(64))
new_ptr_val = Gen.builder.add(ptr_as_int, ir.Constant(ir.IntType(64), byte_offset))
if elem_type == 'i32':
elem_ptr = 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 == 'float':
return Gen.builder.bitcast(alloc_ptr, ir.PointerType(ir.FloatType()), name="list_float_ptr")
elif elem_type == 'double':
return Gen.builder.bitcast(alloc_ptr, ir.PointerType(ir.DoubleType()), name="list_double_ptr")
return cast_ptr
def _HandleConstantLlvm(self, Node, VarType=None):
Gen = self.Trans.LlvmGen
Value = 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):
from lib.includes.t import CTypeRegistry as _CR
resolved = _CR.ResolveName(VarType) or _CR.CNameToClass(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)
is_unsigned = getattr(inst, 'IsSigned', True) is False
mask = (1 << size) - 1
return ir.Constant(ir.IntType(size), Value & mask)
llvm_match = __import__('re').match(r'^i(\d+)$', VarType)
if llvm_match:
width = 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 = getattr(Node, 'kind', None) == 'u'
if len(Value) == 1 and not is_wide:
if VarType is not None and isinstance(VarType, ir.PointerType):
pass
elif isinstance(VarType, ir.IntType):
return ir.Constant(VarType, ord(Value[0]))
elif isinstance(VarType, str):
from lib.includes.t import CTypeRegistry as _CR
resolved = _CR.ResolveName(VarType) or _CR.CNameToClass(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 = __import__('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]))
if is_wide:
str_value = Value + '\x00'
wide_chars = [ord(c) for c in str_value]
target_count = len(wide_chars)
str_type = ir.ArrayType(ir.IntType(16), 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, wide_chars)
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
else:
str_value = Value + '\x00'
str_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, VarType=None):
Gen = self.Trans.LlvmGen
VarName = Node.id
define_constants = getattr(Gen, '_define_constants', {})
if VarName in define_constants:
val = 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):
str_value = val + '\x00'
str_bytes = str_value.encode('utf-8')
str_type = ir.ArrayType(ir.IntType(8), len(str_bytes))
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
# 如果 _define_constants 中没有,尝试从符号表查找
if VarName not in getattr(Gen, '_define_constants', {}):
try:
sym_info = self.translator.SymbolTable.get(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):
str_value = val + '\x00'
str_bytes = str_value.encode('utf-8')
str_type = ir.ArrayType(ir.IntType(8), len(str_bytes))
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
except Exception as _e:
if __import__('lib.constants.config', fromlist=['mode']).mode == "strict":
self.Trans.LogWarning(f"异常被忽略: {_e}")
# 优先检查 variablesalloca因为变量可能已被 += 等操作迁移到 variables
if VarName in Gen.variables and Gen.variables[VarName] is not None:
VarPtr = 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
return Gen._load(VarPtr, name=VarName)
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 = 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 = 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)
if VarName in self.Trans.SymbolTable:
SymInfo = self.Trans.SymbolTable[VarName]
if SymInfo.IsEnumMember and isinstance(SymInfo.value, int):
return ir.Constant(ir.IntType(32), SymInfo.value)
if SymInfo.IsFunction or SymInfo.IsFuncPtr:
MangledName = Gen._mangle_func_name(VarName)
if MangledName in Gen.module.globals:
g = 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 = 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 = getattr(self.Trans, '_t_c_imported_names', {})
if VarName in t_c_imported:
src_module, src_name = t_c_imported[VarName]
virtual_attr = 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)