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

452 lines
25 KiB
Python
Raw Permalink 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, 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}")
# 优先检查 variablesalloca因为变量可能已被 += 等操作迁移到 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
# 防御性恢复:对于 double-pointer如 ndarray**pointee 是 PointerType
# 或其他非结构体类型,若 global_struct_class 已记录该变量的类名,
# 也恢复 var_struct_class。这确保方法调用能查找到正确的特化类
# (如 ndarray[double]),而非 fallback 到 opaque 基础类型。
if VarName in Gen.global_struct_class and VarName not in Gen.var_struct_class:
Gen.var_struct_class[VarName] = Gen.global_struct_class[VarName]
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)