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

968 lines
51 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, CTypeInfo, FuncMeta
from lib.includes import t
import ast
import llvmlite.ir as ir
class ClassHandle(BaseHandle):
def _is_exception_class(self, Node):
if not Node.bases:
return False
for base in Node.bases:
if hasattr(base, 'id'):
if base.id == 'Exception':
return True
if base.id in self.Trans.exception_registry:
return True
elif hasattr(base, 'attr'):
if base.attr == 'Exception':
return True
if base.attr in self.Trans.exception_registry:
return True
return False
def _get_exception_parent(self, Node):
if not Node.bases:
return None
for base in Node.bases:
if hasattr(base, 'id'):
if base.id != 'Exception' and base.id in self.Trans.exception_registry:
return base.id
elif hasattr(base, 'attr'):
if base.attr != 'Exception' and base.attr in self.Trans.exception_registry:
return base.attr
return None
def _RegisterExceptionClass(self, Node):
ClassName = Node.name
if ClassName in self.Trans.exception_registry:
return
code = self.Trans._next_exception_code
self.Trans._next_exception_code += 1
self.Trans.exception_registry[ClassName] = code
parent = self._get_exception_parent(Node)
if parent:
self.Trans.exception_parents[ClassName] = parent
ExcTypeInfo = CTypeInfo()
ExcTypeInfo.Name = ClassName
ExcTypeInfo.IsExceptionClass = True
ExcTypeInfo.value = code
self.Trans.SymbolTable.insert(ClassName, ExcTypeInfo)
def _is_generic_class(self, Node):
if hasattr(Node, 'type_params') and Node.type_params:
return True
return False
def _mangle_generic_class_name(self, class_name, type_args):
from lib.includes.t import CTypeRegistry
mangled_args = []
for ta in type_args:
llvm_str = CTypeRegistry.NameToLLVM(ta)
if llvm_str:
if llvm_str in ('float', 'double', 'half', 'fp128'):
mangled = 'f' + str(CTypeRegistry.GetClassByName(ta)().Size)
else:
mangled = llvm_str
else:
mangled = ta
mangled_args.append(mangled)
return class_name + '[' + ']['.join(mangled_args) + ']'
def _specialize_generic_class(self, ClassName, type_args, Gen, type_names=None):
if not hasattr(self, '_generic_class_templates') or ClassName not in self._generic_class_templates:
return None
template = self._generic_class_templates[ClassName]
Node = template['node']
type_param_names = template['type_params']
if len(type_args) != len(type_param_names):
return None
spec_key = ClassName + '<' + ','.join(type_args) + '>'
if not hasattr(self, '_generic_class_specializations'):
self._generic_class_specializations = {}
if spec_key in self._generic_class_specializations:
return self._generic_class_specializations[spec_key]
spec_name = self._mangle_generic_class_name(ClassName, type_args)
type_map = {}
for i, tp_name in enumerate(type_param_names):
if type_names and i < len(type_names):
type_map[tp_name] = type_names[i]
else:
type_map[tp_name] = type_args[i]
if type_names:
if not hasattr(self.Trans, '_t_c_imported_names'):
self.Trans._t_c_imported_names = {}
for tn in type_names:
if tn not in self.Trans._t_c_imported_names:
self.Trans._t_c_imported_names[tn] = ('t', tn)
import copy
SpecNode = copy.deepcopy(Node)
SpecNode.type_params = []
SpecNode.name = spec_name
for item in SpecNode.body:
if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name):
if item.annotation:
item.annotation = self.Trans.FunctionHandler._replace_type_in_annotation(item.annotation, type_map)
elif isinstance(item, ast.FunctionDef):
for arg in item.args.args:
if arg.annotation:
arg.annotation = self.Trans.FunctionHandler._replace_type_in_annotation(arg.annotation, type_map)
if item.returns:
item.returns = self.Trans.FunctionHandler._replace_type_in_annotation(item.returns, type_map)
self.Trans.FunctionHandler._apply_type_map_to_body(item.body, type_map)
saved_builder = Gen.builder
saved_func = Gen.func
saved_variables = dict(Gen.variables) if Gen.variables else {}
saved_direct_values = dict(Gen._direct_values) if Gen._direct_values else {}
saved_var_type_info = dict(Gen.var_type_info) if Gen.var_type_info else {}
saved_var_signedness = dict(Gen.var_signedness) if Gen.var_signedness else {}
saved_global_vars = set(Gen.global_vars) if Gen.global_vars else set()
saved_var_scopes = [dict(s) for s in self.Trans.VarScopes] if self.Trans.VarScopes else []
saved_block = None
if Gen.builder and Gen.builder.block and not Gen.builder.block.is_terminated:
saved_block = Gen.builder.block
self._EmitClassLlvm(SpecNode, Gen)
Gen.builder = saved_builder
if saved_block is not None and Gen.builder is not None:
Gen.builder.position_at_end(saved_block)
Gen.func = saved_func
Gen.variables = saved_variables
Gen._direct_values = saved_direct_values
Gen.var_type_info = saved_var_type_info
Gen.var_signedness = saved_var_signedness
Gen.global_vars = saved_global_vars
self.Trans.VarScopes = saved_var_scopes
self._generic_class_specializations[spec_key] = spec_name
if not hasattr(self.Trans, '_generic_class_specializations'):
self.Trans._generic_class_specializations = {}
self.Trans._generic_class_specializations[spec_key] = spec_name
if hasattr(self.Trans, '_module_sha1') and self.Trans._module_sha1:
Gen._struct_sha1_map[spec_name] = self.Trans._module_sha1
return spec_name
def _EmitClassLlvm(self, Node, Gen):
ClassName = Node.name
if self._is_generic_class(Node):
if not hasattr(self, '_generic_class_templates'):
self._generic_class_templates = {}
type_params = [tp.name for tp in Node.type_params]
self._generic_class_templates[ClassName] = {
'node': Node,
'type_params': type_params,
}
return
if self._is_exception_class(Node):
self._RegisterExceptionClass(Node)
return
IsCenum = False
IsCunion = False
IsRenum = False
if Node.bases:
for base in Node.bases:
if hasattr(base, 'attr'):
if base.attr == 'CEnum' or base.attr == 'Enum':
IsCenum = True
break
elif base.attr == 'CUnion':
IsCunion = True
elif base.attr == 'REnum':
IsRenum = True
elif hasattr(base, 'id'):
if base.id == 'CEnum' or base.id == 'Enum':
IsCenum = True
break
elif base.id == 'CUnion':
IsCunion = True
elif base.id == 'REnum':
IsRenum = True
if IsCenum:
self._RegisterEnumMembers(Node)
return
if IsCunion:
self._EmitUnionLlvm(Node, Gen)
return
if IsRenum:
self._EmitREnumLlvm(Node, Gen)
return
IsCpythonObject = False
IsCVTable = False
if hasattr(Node, 'decorator_list') and Node.decorator_list:
for decorator in Node.decorator_list:
if isinstance(decorator, ast.Attribute):
if hasattr(decorator.value, 'id') and decorator.value.id == 't':
if decorator.attr == 'Object':
IsCpythonObject = True
elif decorator.attr == 'CVTable':
IsCVTable = True
elif isinstance(decorator, ast.Name):
if decorator.id == 'Object':
IsCpythonObject = True
elif decorator.id == 'CVTable':
IsCVTable = True
elif isinstance(decorator, ast.Call):
# 检测 @c.Attribute(t.attr.packed)
if isinstance(decorator.func, ast.Attribute):
if getattr(decorator.func.value, 'id', None) == 'c' and decorator.func.attr == 'Attribute':
for arg in decorator.args:
if isinstance(arg, ast.Attribute):
if isinstance(arg.value, ast.Attribute):
if getattr(arg.value.value, 'id', None) == 't' and arg.value.attr == 'attr' and arg.attr == 'packed':
Gen.class_packed.add(ClassName)
HasMethods = any(isinstance(item, ast.FunctionDef) for item in Node.body)
if HasMethods and not IsCpythonObject:
IsCpythonObject = True
HasParentClass = False
ParentClassName = None
if Node.bases:
for base in Node.bases:
base_name = None
if hasattr(base, 'id'):
base_name = base.id
elif hasattr(base, 'attr'):
base_name = base.attr
if base_name and base_name not in ('CEnum', 'Enum', 'CUnion', 'CStruct', 'Object', 'CVTable', 'Exception'):
if base_name in Gen.class_members or base_name in Gen.structs:
HasParentClass = True
ParentClassName = base_name
break
if HasParentClass and not IsCVTable:
IsCVTable = True
if IsCVTable:
Gen.class_vtable.add(ClassName)
if ParentClassName:
Gen.class_vtable.add(ParentClassName)
if ClassName not in Gen.class_methods:
Gen.class_methods[ClassName] = []
# 当前文件定义的类是权威定义,覆盖之前导入阶段可能添加的外部声明
Gen.class_members[ClassName] = []
Gen.class_member_defaults[ClassName] = {}
Gen.class_member_signeds[ClassName] = {}
Gen.class_member_bitfields[ClassName] = {}
Gen.class_member_byteorders[ClassName] = {}
Gen.class_member_bitoffsets[ClassName] = {}
ParentClass = None
if Node.bases:
for base in Node.bases:
if hasattr(base, 'id'):
base_name = base.id
if base_name in Gen.class_members and base_name != ClassName:
ParentClass = base_name
break
elif hasattr(base, 'attr'):
base_name = base.attr
if base_name in Gen.class_members and base_name != ClassName:
ParentClass = base_name
break
if ParentClass:
Gen.class_parent[ClassName] = ParentClass
if ParentClass:
parent_members = list(Gen.class_members[ParentClass])
parent_defaults = dict(Gen.class_member_defaults.get(ParentClass, {}))
existing_names = {m[0] for m in Gen.class_members[ClassName]}
inherited = []
for pm_name, pm_type in parent_members:
if pm_name not in existing_names:
inherited.append((pm_name, pm_type))
if pm_name in parent_defaults:
Gen.class_member_defaults[ClassName][pm_name] = parent_defaults[pm_name]
Gen.class_members[ClassName] = inherited + Gen.class_members[ClassName]
if ParentClass in Gen.class_methods:
parent_methods = list(Gen.class_methods[ParentClass])
for pm in parent_methods:
method_name = pm.split('.')[-1] if '.' in pm else pm
child_method = f"{ClassName}.{method_name}"
if child_method not in Gen.class_methods[ClassName]:
Gen.class_methods[ClassName].append(child_method)
Gen.register_method(ClassName, child_method)
for item in Node.body:
if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name):
VarName = item.target.id
try:
TypeInfo = CTypeInfo.FromNode(item.annotation, self.Trans.SymbolTable)
if TypeInfo is None:
TypeInfo = CTypeInfo()
TypeInfo.BaseType = t.CInt()
IsPtr = TypeInfo.IsPtr
MemberType = Gen._ctype_to_llvm(TypeInfo)
if isinstance(MemberType, ir.VoidType):
MemberType = ir.PointerType(ir.IntType(8))
Gen.class_members[ClassName].append((VarName, MemberType))
if isinstance(MemberType, ir.PointerType) and isinstance(MemberType.pointee, ir.IntType) and MemberType.pointee.width == 8:
if isinstance(item.annotation, ast.BinOp) and isinstance(item.annotation.op, ast.BitOr):
def _find_struct_in_annot(node):
if isinstance(node, ast.Name) and node.id in Gen.structs:
return node.id
if isinstance(node, ast.Constant) and isinstance(node.value, str) and node.value in Gen.structs:
return node.value
if isinstance(node, ast.Attribute) and node.attr in Gen.structs:
return node.attr
if isinstance(node, ast.BinOp) and isinstance(node.op, ast.BitOr):
result = _find_struct_in_annot(node.left)
if result:
return result
return _find_struct_in_annot(node.right)
return None
element_class = _find_struct_in_annot(item.annotation)
if element_class:
if ClassName not in Gen.class_member_element_class:
Gen.class_member_element_class[ClassName] = {}
Gen.class_member_element_class[ClassName][VarName] = element_class
if TypeInfo and hasattr(TypeInfo.BaseType, 'IsSigned'):
Gen.class_member_signeds[ClassName][VarName] = TypeInfo.BaseType.IsSigned
else:
Gen.class_member_signeds[ClassName][VarName] = None
if TypeInfo and TypeInfo.IsBitField:
Gen.class_member_bitfields[ClassName][VarName] = TypeInfo.BitWidth
else:
Gen.class_member_bitfields[ClassName][VarName] = 0
if TypeInfo and TypeInfo.ByteOrder:
Gen.class_member_byteorders[ClassName][VarName] = TypeInfo.ByteOrder
else:
Gen.class_member_byteorders[ClassName][VarName] = ""
if item.value:
const = self._BuildScalarConstant(item.value, MemberType)
if const:
Gen.class_member_defaults[ClassName][VarName] = const
except Exception: # 回退:类成员类型解析失败时使用默认 i32
Gen.class_members[ClassName].append((VarName, ir.IntType(32)))
Gen.class_member_signeds[ClassName][VarName] = None
Gen.class_member_bitfields[ClassName][VarName] = 0
existing_members = {m[0] for m in Gen.class_members[ClassName]}
for item in Node.body:
if isinstance(item, ast.FunctionDef):
self._current_method_args = {}
for arg in item.args.args:
if arg.arg != 'self' and arg.annotation:
self._current_method_args[arg.arg] = arg.annotation
self._scan_method_body_for_self_members(item.body, ClassName, existing_members, Gen)
self._current_method_args = {}
for item in Node.body:
if isinstance(item, ast.FunctionDef):
MethodName = item.name
FullMethodName = f"{ClassName}.{MethodName}"
# property setter/deleter 使用不同的函数名后缀
item_meta = self.Trans.FunctionHandler._ExtractFuncMeta(item.decorator_list)
if FuncMeta.PROPERTY_SETTER in item_meta:
FullMethodName = FullMethodName + '$set'
elif FuncMeta.PROPERTY_DELETER in item_meta:
FullMethodName = FullMethodName + '$del'
if FullMethodName not in Gen.class_methods[ClassName]:
Gen.class_methods[ClassName].append(FullMethodName)
Gen.register_method(ClassName, FullMethodName)
if IsCVTable:
Gen.class_vtable.add(ClassName)
Gen._generate_structs()
Gen._create_Vtable_globals()
if (IsCpythonObject or IsCVTable) and ClassName in Gen.structs:
self._EmitNewFunctionLlvm(ClassName, Gen, IsCVTable=IsCVTable, IsCpythonObject=IsCpythonObject)
for item in Node.body:
if isinstance(item, ast.FunctionDef):
self.Trans.FunctionHandler._EmitFunctionForwardDeclLlvm(item, Gen, ClassName=ClassName)
for item in Node.body:
if isinstance(item, ast.FunctionDef):
self.Trans.FunctionHandler._EmitFunctionLlvm(item, Gen, ClassName=ClassName)
if IsCVTable and ClassName in Gen.Vtables:
self._fill_vtable_with_methods(ClassName, Gen)
# Generate wrapper functions for inherited methods that are not overridden
if ParentClass:
self._generate_inherited_method_wrappers(ClassName, Gen)
def _generate_inherited_method_wrappers(self, ClassName, Gen):
"""为子类继承但未覆写的方法生成包装函数,使跨模块调用能正确链接。
例如 Label 继承 Widget.place生成 Label.place 函数,
内部将 self 从 Label* bitcast 为 Widget*,然后调用 Widget.place。
"""
methods = Gen.class_methods.get(ClassName, [])
if not methods:
return
if ClassName not in Gen.structs:
return
for method_name in methods:
method_short = method_name.split('.')[-1] if '.' in method_name else method_name
child_full = f"{ClassName}.{method_short}"
# 如果子类已经有该方法的实现,跳过
if Gen._find_function(child_full):
continue
# 沿继承链查找父类实现
parent_func = None
parent_class = None
p = Gen.class_parent.get(ClassName)
while p:
parent_full = f"{p}.{method_short}"
parent_func = Gen._find_function(parent_full)
if parent_func:
parent_class = p
break
p = Gen.class_parent.get(p)
if not parent_func or not parent_class:
continue
# 父类和子类都需要有 struct 定义
if parent_class not in Gen.structs:
continue
# 构造子类包装函数签名:与父类函数相同,但 self 参数类型为子类指针
parent_ftype = parent_func.function_type
parent_ret_type = parent_ftype.return_type
parent_param_types = list(parent_ftype.args)
is_vararg = parent_ftype.var_arg
# 判断父类方法是否是静态方法:
# 如果第一个参数类型是父类指针,说明是实例方法;否则是静态方法
parent_struct_ptr_type = ir.PointerType(Gen.structs[parent_class]) if parent_class in Gen.structs else None
is_static = True
if parent_param_types and parent_struct_ptr_type:
first_param = parent_param_types[0]
# 检查第一个参数是否是父类指针(或可以 bitcast 为父类指针的指针类型)
if isinstance(first_param, ir.PointerType):
if isinstance(first_param.pointee, ir.IdentifiedStructType):
if first_param.pointee.name == Gen.structs[parent_class].name:
is_static = False
elif first_param == ir.IntType(8):
# i8* 可能是前向声明时的 self 类型,视为实例方法
is_static = False
if is_static:
# 静态方法:签名与父类完全一致,不需要 self 参数
child_param_types = list(parent_param_types)
child_struct_ptr = None
else:
# 实例方法替换第一个参数self的类型为子类指针
child_struct_ptr = ir.PointerType(Gen.structs[ClassName])
if parent_param_types:
child_param_types = [child_struct_ptr] + parent_param_types[1:]
else:
child_param_types = [child_struct_ptr]
child_func_type = ir.FunctionType(parent_ret_type, child_param_types, var_arg=is_vararg)
child_mangled = Gen._mangle_func_name(child_full)
child_func = Gen._get_or_declare_function(child_mangled, child_func_type)
Gen.functions[child_full] = child_func
# 生成函数体
entry_block = child_func.append_basic_block(name="entry")
saved_builder = Gen.builder
saved_func = Gen.func
Gen.builder = ir.IRBuilder(entry_block)
Gen.func = child_func
# 准备参数
call_args = []
if is_static:
# 静态方法:参数直接传递,不做 bitcast
for arg in child_func.args:
call_args.append(arg)
else:
# 实例方法bitcast self其余参数直接传递
for i, arg in enumerate(child_func.args):
if i == 0 and parent_param_types:
actual_self_type = parent_param_types[0]
casted_self = Gen.builder.bitcast(arg, actual_self_type, name="self_cast")
call_args.append(casted_self)
else:
call_args.append(arg)
result = Gen.builder.call(parent_func, call_args, name="inherited_call")
if isinstance(parent_ret_type, ir.VoidType):
Gen.builder.ret_void()
else:
Gen.builder.ret(result)
Gen.builder = saved_builder
Gen.func = saved_func
def _fill_vtable_with_methods(self, ClassName, Gen):
methods = Gen.class_methods.get(ClassName, [])
if not methods or ClassName not in Gen.Vtables:
return
Vtable = Gen.Vtables[ClassName]
VtableType = Vtable.type.pointee
null_i8ptr = ir.Constant(ir.PointerType(ir.IntType(8)), None)
initializers = []
for mi, method_name in enumerate(methods):
method_short = method_name.split('.')[-1] if '.' in method_name else method_name
func = Gen._find_function(f"{ClassName}.{method_short}")
if not func:
func = Gen._find_function(f"{ClassName}.{method_short}__")
if not func:
p = Gen.class_parent.get(ClassName)
while p:
func = Gen._find_function(f"{p}.{method_short}")
if not func:
func = Gen._find_function(f"{p}.{method_short}__")
if func:
break
p = Gen.class_parent.get(p)
if func:
initializers.append(func.bitcast(ir.PointerType(ir.IntType(8))))
else:
initializers.append(null_i8ptr)
if initializers:
Vtable.initializer = ir.Constant(VtableType, initializers)
def _scan_method_body_for_self_members(self, body, ClassName, existing_members, Gen):
for stmt in body:
self._scan_node_for_self_assign(stmt, ClassName, existing_members, Gen)
def _scan_node_for_self_assign(self, node, ClassName, existing_members, Gen):
if isinstance(node, ast.AnnAssign):
if (isinstance(node.target, ast.Attribute)
and isinstance(node.target.value, ast.Name)
and node.target.value.id == 'self'):
VarName = node.target.attr
if VarName not in existing_members:
try:
TypeInfo = CTypeInfo.FromNode(node.annotation, self.Trans.SymbolTable)
if TypeInfo is None:
TypeInfo = CTypeInfo()
TypeInfo.BaseType = t.CInt()
MemberType = Gen._ctype_to_llvm(TypeInfo)
if isinstance(MemberType, ir.VoidType):
MemberType = ir.IntType(32)
except Exception: # 回退:成员类型解析失败时使用默认 i32
MemberType = ir.IntType(32)
Gen.class_members[ClassName].append((VarName, MemberType))
Gen.class_member_signeds[ClassName][VarName] = None
Gen.class_member_bitfields[ClassName][VarName] = 0
Gen.class_member_byteorders[ClassName][VarName] = ""
Gen.class_member_bitoffsets[ClassName][VarName] = 0
existing_members.add(VarName)
len_name = f'{VarName}__len'
if isinstance(MemberType, ir.PointerType) and node.value and isinstance(node.value, ast.List):
if len_name not in existing_members:
Gen.class_members[ClassName].append((len_name, ir.IntType(64)))
Gen.class_member_signeds[ClassName][len_name] = None
Gen.class_member_bitfields[ClassName][len_name] = 0
Gen.class_member_byteorders[ClassName][len_name] = ""
Gen.class_member_bitoffsets[ClassName][len_name] = 0
Gen.class_member_defaults[ClassName][len_name] = ir.Constant(ir.IntType(64), len(node.value.elts))
existing_members.add(len_name)
elif isinstance(node, ast.Assign):
for target in node.targets:
if (isinstance(target, ast.Attribute)
and isinstance(target.value, ast.Name)
and target.value.id == 'self'):
VarName = target.attr
if VarName not in existing_members:
MemberType = ir.IntType(32)
if node.value:
try:
if isinstance(node.value, ast.Constant):
if isinstance(node.value.value, int):
MemberType = ir.IntType(32)
elif isinstance(node.value.value, float):
MemberType = ir.DoubleType()
elif isinstance(node.value.value, str):
MemberType = ir.PointerType(ir.IntType(8))
elif isinstance(node.value, ast.Name):
if hasattr(self, '_current_method_args'):
arg_info = self._current_method_args.get(node.value.id)
if arg_info:
TypeInfo = CTypeInfo.FromNode(arg_info, self.Trans.SymbolTable)
if TypeInfo:
InferredType = Gen._ctype_to_llvm(TypeInfo)
if not isinstance(InferredType, ir.VoidType):
MemberType = InferredType
elif isinstance(node.value, ast.Call):
if isinstance(node.value.func, ast.Name):
if node.value.func.id in ('float', 'double'):
MemberType = ir.DoubleType()
elif node.value.func.id in ('int', 'long', 'short', 'char'):
MemberType = ir.IntType(32)
elif isinstance(node.value.func, ast.Attribute):
if node.value.func.attr == 'len':
MemberType = ir.IntType(32)
except Exception as _e:
if __import__('lib.constants.config', fromlist=['mode']).mode == "strict":
self.Trans.LogWarning(f"异常被忽略: {_e}")
Gen.class_members[ClassName].append((VarName, MemberType))
Gen.class_member_signeds[ClassName][VarName] = None
Gen.class_member_bitfields[ClassName][VarName] = 0
Gen.class_member_byteorders[ClassName][VarName] = ""
Gen.class_member_bitoffsets[ClassName][VarName] = 0
existing_members.add(VarName)
for child in ast.iter_child_nodes(node):
if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)):
continue
self._scan_node_for_self_assign(child, ClassName, existing_members, Gen)
def _EmitUnionLlvm(self, Node, Gen):
ClassName = Node.name
union_member_types = []
for item in Node.body:
if isinstance(item, ast.ClassDef):
nested_class_name = item.name
nested_member_types = []
for nested_item in item.body:
if isinstance(nested_item, ast.AnnAssign) and isinstance(nested_item.target, ast.Name):
VarName = nested_item.target.id
try:
TypeInfo = CTypeInfo.FromNode(nested_item.annotation, self.Trans.SymbolTable)
if TypeInfo is None:
TypeInfo = CTypeInfo()
TypeInfo.BaseType = t.CInt()
MemberType = Gen._ctype_to_llvm(TypeInfo)
if isinstance(MemberType, ir.VoidType):
MemberType = ir.PointerType(ir.IntType(8))
nested_member_types.append(MemberType)
except Exception: # 回退:嵌套成员类型解析失败时使用默认 i32
nested_member_types.append(ir.IntType(32))
if nested_member_types:
nested_struct_type = ir.IdentifiedStructType(Gen.module, f"{ClassName}_{nested_class_name}")
nested_struct_type.set_body(*nested_member_types)
Gen.structs[f"{ClassName}_{nested_class_name}"] = nested_struct_type
union_member_types.append(nested_struct_type)
nested_class_full_name = f"{ClassName}_{nested_class_name}"
if nested_class_full_name not in Gen.class_members:
Gen.class_members[nested_class_full_name] = []
if nested_class_full_name not in Gen.class_member_signeds:
Gen.class_member_signeds[nested_class_full_name] = {}
for nested_item in item.body:
if isinstance(nested_item, ast.AnnAssign) and isinstance(nested_item.target, ast.Name):
VarName = nested_item.target.id
try:
TypeInfo = CTypeInfo.FromNode(nested_item.annotation, self.Trans.SymbolTable)
if TypeInfo is None:
TypeInfo = CTypeInfo()
TypeInfo.BaseType = t.CInt()
MemberType = Gen._ctype_to_llvm(TypeInfo)
if isinstance(MemberType, ir.VoidType):
MemberType = ir.PointerType(ir.IntType(8))
Gen.class_members[nested_class_full_name].append((VarName, MemberType))
if TypeInfo and hasattr(TypeInfo.BaseType, 'IsSigned'):
Gen.class_member_signeds[nested_class_full_name][VarName] = TypeInfo.BaseType.IsSigned
else:
Gen.class_member_signeds[nested_class_full_name][VarName] = None
except Exception: # 回退:嵌套类成员类型解析失败时使用默认 i32
Gen.class_members[nested_class_full_name].append((VarName, ir.IntType(32)))
Gen.class_member_signeds[nested_class_full_name][VarName] = None
if union_member_types:
max_size = 0
for member_type in union_member_types:
try:
size = member_type.get_abi_size(ir.DataLayout(Gen.module.data_layout))
except Exception: # 回退:获取类型大小失败时使用默认值 8
size = 8
max_size = max(max_size, size)
union_type = ir.IdentifiedStructType(Gen.module, ClassName)
union_type.set_body(ir.ArrayType(ir.IntType(8), max_size))
Gen.structs[ClassName] = union_type
UnionNode = CTypeInfo()
UnionNode.Name = ClassName
UnionNode.IsUnion = True
self.Trans.SymbolTable.insert(ClassName, UnionNode)
def _EmitREnumLlvm(self, Node, Gen):
ClassName = Node.name
variant_info = []
variant_names = []
variant_index = 0
for item in Node.body:
if isinstance(item, ast.ClassDef):
VariantName = item.name
variant_names.append(VariantName)
member_types = []
member_names = []
member_annotations = []
for nested_item in item.body:
if isinstance(nested_item, ast.AnnAssign) and isinstance(nested_item.target, ast.Name):
VarName = nested_item.target.id
try:
TypeInfo = CTypeInfo.FromNode(nested_item.annotation, self.Trans.SymbolTable)
if TypeInfo is None:
TypeInfo = CTypeInfo()
TypeInfo.BaseType = t.CInt()
MemberType = Gen._ctype_to_llvm(TypeInfo)
if isinstance(MemberType, ir.VoidType):
MemberType = ir.PointerType(ir.IntType(8))
member_types.append(MemberType)
member_names.append(VarName)
member_annotations.append(nested_item.annotation)
except Exception: # 回退:成员类型解析失败时使用默认 i32
member_types.append(ir.IntType(32))
member_names.append(VarName)
member_annotations.append(None)
variant_info.append((VariantName, member_types, member_names, member_annotations, item.lineno, variant_index))
variant_index += 1
elif isinstance(item, ast.Assign):
for target in item.targets:
if isinstance(target, ast.Name):
VarName = target.id
variant_names.append(VarName)
value = None
if item.value and isinstance(item.value, ast.Constant):
value = item.value.value
variant_index = value + 1
else:
value = variant_index
variant_index += 1
variant_info.append((VarName, [], [], [], item.lineno, value))
max_variant_struct = None
max_variant_size = 0
for _, member_types, _, _, _, _ in variant_info:
if member_types:
full_types = [ir.IntType(32)] + member_types
size = 4
for mt in full_types:
if isinstance(mt, ir.IntType):
size += mt.width // 8
elif isinstance(mt, ir.PointerType):
size += 8
elif isinstance(mt, ir.FloatType):
size += 4
elif isinstance(mt, ir.DoubleType):
size += 8
elif isinstance(mt, ir.ArrayType):
if isinstance(mt.element, ir.IntType):
size += mt.element.width // 8 * mt.count
else:
size += 8 * mt.count
else:
size += 8
if size > max_variant_size:
max_variant_size = size
max_variant_struct = full_types
if max_variant_struct is None:
max_variant_struct = [ir.IntType(32), ir.IntType(32)]
renum_type = ir.IdentifiedStructType(Gen.module, ClassName)
renum_type.set_body(*max_variant_struct)
Gen.structs[ClassName] = renum_type
for VariantName, member_types, member_names, member_annotations, lineno, tag_value in variant_info:
NestedStructName = f"{ClassName}_{VariantName}"
if member_types:
padded_member_types = list(max_variant_struct)
nested_struct_type = ir.IdentifiedStructType(Gen.module, NestedStructName)
nested_struct_type.set_body(*padded_member_types)
Gen.structs[NestedStructName] = nested_struct_type
if NestedStructName not in Gen.class_members:
Gen.class_members[NestedStructName] = []
if NestedStructName not in Gen.class_member_signeds:
Gen.class_member_signeds[NestedStructName] = {}
Gen.class_members[NestedStructName].append(('__tag', ir.IntType(32)))
Gen.class_member_signeds[NestedStructName]['__tag'] = None
for i, (mname, mtype) in enumerate(zip(member_names, member_types)):
Gen.class_members[NestedStructName].append((mname, mtype))
try:
ti = CTypeInfo.FromNode(member_annotations[i], self.Trans.SymbolTable)
if ti and hasattr(ti.BaseType, 'IsSigned'):
Gen.class_member_signeds[NestedStructName][mname] = ti.BaseType.IsSigned
else:
Gen.class_member_signeds[NestedStructName][mname] = None
except Exception: # 回退:签名信息解析失败时设为 None
Gen.class_member_signeds[NestedStructName][mname] = None
MemberNode = CTypeInfo()
MemberNode.Name = VariantName
MemberNode.BaseType = t.CEnum(ClassName)
MemberNode.value = tag_value
MemberNode.EnumName = ClassName
MemberNode.Lineno = lineno
MemberNode.IsEnumMember = True
self.Trans.SymbolTable.insert(VariantName, MemberNode)
RenumTypeInfo = CTypeInfo()
RenumTypeInfo.Name = ClassName
RenumTypeInfo.BaseType = t.REnum(ClassName)
RenumTypeInfo.IsRenum = True
RenumTypeInfo.IsEnum = True
RenumTypeInfo.RenumVariants = variant_names
self.Trans.SymbolTable.insert(ClassName, RenumTypeInfo)
def _RegisterEnumMembers(self, Node):
ClassName = Node.name
EnumTypeInfo = CTypeInfo()
EnumTypeInfo.Name = ClassName
EnumTypeInfo.BaseType = t.CEnum(ClassName)
EnumTypeInfo.IsEnum = True
self.Trans.SymbolTable.insert(ClassName, EnumTypeInfo)
from lib.core.Exportable import EnumMember as ExportEnumMember
enum_export = self.Trans.Exportable.add_enum(
name=ClassName,
lineno=Node.lineno,
is_public=True
)
next_enum_value = 0
for item in Node.body:
if isinstance(item, ast.Assign):
for target in item.targets:
if isinstance(target, ast.Name):
VarName = target.id
value = None
if item.value:
if isinstance(item.value, ast.Constant):
value = item.value.value
next_enum_value = value + 1
elif isinstance(item.value, ast.UnaryOp) and isinstance(item.value.op, ast.USub):
if isinstance(item.value.operand, ast.Constant):
value = -item.value.operand.value
next_enum_value = value + 1
elif isinstance(item.value, ast.Name):
value = item.value.id
if self.Trans.SymbolTable.has(value):
ref_info = self.Trans.SymbolTable[value]
if ref_info.value is not None and isinstance(ref_info.value, int):
next_enum_value = ref_info.value + 1
else:
value = next_enum_value
next_enum_value += 1
MemberNode = CTypeInfo()
MemberNode.Name = VarName
MemberNode.BaseType = t.CEnum(ClassName)
MemberNode.value = value
MemberNode.EnumName = ClassName
MemberNode.Lineno = item.lineno
MemberNode.IsEnumMember = True
self.Trans.SymbolTable.insert(VarName, MemberNode)
enum_export.members.append(ExportEnumMember(
name=VarName,
value=value,
lineno=item.lineno
))
elif isinstance(item, ast.AnnAssign):
if isinstance(item.target, ast.Name):
VarName = item.target.id
value = None
if item.value:
if isinstance(item.value, ast.Constant):
value = item.value.value
next_enum_value = value + 1
elif isinstance(item.value, ast.UnaryOp) and isinstance(item.value.op, ast.USub):
if isinstance(item.value.operand, ast.Constant):
value = -item.value.operand.value
next_enum_value = value + 1
elif isinstance(item.value, ast.Name):
value = item.value.id
if self.Trans.SymbolTable.has(value):
ref_info = self.Trans.SymbolTable[value]
if ref_info.value is not None and isinstance(ref_info.value, int):
next_enum_value = ref_info.value + 1
else:
value = next_enum_value
next_enum_value += 1
MemberNode = CTypeInfo()
MemberNode.Name = VarName
MemberNode.BaseType = t.CEnum(ClassName)
MemberNode.value = value
MemberNode.EnumName = ClassName
MemberNode.Lineno = item.lineno
MemberNode.IsEnumMember = True
self.Trans.SymbolTable.insert(VarName, MemberNode)
enum_export.members.append(ExportEnumMember(
name=VarName,
value=value,
lineno=item.lineno
))
def _EmitNewFunctionLlvm(self, ClassName, Gen, IsCVTable=False, IsCpythonObject=False):
NewFuncName = f'{ClassName}.__before_init__'
existing_func = Gen.functions.get(NewFuncName)
if existing_func and len(existing_func.blocks) > 0:
return
MangledName = Gen._mangle_name(NewFuncName)
StructType = Gen.structs[ClassName]
if isinstance(StructType, ir.IdentifiedStructType) and (StructType.elements is None or len(StructType.elements) == 0):
self.Trans.ImportHandler._TryLoadStructFromStub(ClassName, Gen)
StructType = Gen.structs.get(ClassName, StructType)
StructPtrType = ir.PointerType(StructType)
FuncType = ir.FunctionType(ir.VoidType(), [StructPtrType])
if existing_func:
func = existing_func
else:
func = Gen._get_or_declare_function(MangledName, FuncType)
Gen.functions[NewFuncName] = func
EntryBlock = func.append_basic_block(name="entry")
saved_builder = Gen.builder
saved_func = Gen.func
Gen.builder = ir.IRBuilder(EntryBlock)
Gen.func = func
StructPtr = func.args[0]
if IsCVTable:
base = 1
if ClassName in Gen.Vtables:
VtablePtr = Gen.builder.gep(StructPtr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), 0)], name="vtable_slot")
VtableAddr = Gen.builder.bitcast(Gen.Vtables[ClassName], ir.PointerType(ir.IntType(8)), name="vtable_addr")
Gen.builder.store(VtableAddr, VtablePtr)
members = Gen.class_members.get(ClassName, [])
defaults = Gen.class_member_defaults.get(ClassName, {})
for i, (member_name, member_type) in enumerate(members):
ElemPtr = Gen.builder.gep(StructPtr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), base + i)], name=f"{ClassName}_{member_name}")
if member_name in defaults:
try:
val = defaults[member_name]
if not isinstance(val, ir.Constant):
val = ir.Constant(ElemPtr.type.pointee, val)
Gen.builder.store(val, ElemPtr)
except Exception as _e:
if __import__('lib.constants.config', fromlist=['mode']).mode == "strict":
self.Trans.LogWarning(f"异常被忽略: {_e}")
Gen.builder.ret_void()
Gen.builder = saved_builder
Gen.func = saved_func
def _BuildScalarConstant(self, node, llvm_type):
if isinstance(node, ast.Constant):
if isinstance(node.value, bool):
try:
return ir.Constant(llvm_type, 1 if node.value else 0)
except Exception: # 回退:常量创建失败
return None
elif isinstance(node.value, int):
try:
return ir.Constant(llvm_type, node.value)
except Exception: # 回退:常量创建失败
return None
elif isinstance(node.value, float):
try:
return ir.Constant(llvm_type, node.value)
except Exception: # 回退:常量创建失败
return None
elif isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.USub):
if isinstance(node.operand, ast.Constant) and isinstance(node.operand.value, int):
try:
return ir.Constant(llvm_type, -node.operand.value)
except Exception: # 回退:常量创建失败
return None
return None
def _BuildStructConstant(self, call_node, StructName, Gen):
struct_type = Gen.structs[StructName]
members = Gen.class_members.get(StructName, [])
defaults = Gen.class_member_defaults.get(StructName, {})
member_values = []
for member_name, member_type in members:
if member_name in defaults:
val = defaults[member_name]
if not isinstance(val, ir.Constant):
try:
val = ir.Constant(member_type, val)
except Exception: # 回退:常量转换失败时使用零值
val = ir.Constant(member_type, 0)
member_values.append(val)
else:
try:
if isinstance(member_type, (ir.IdentifiedStructType, ir.LiteralStructType, ir.PointerType, ir.ArrayType)):
member_values.append(ir.Constant(member_type, None))
else:
member_values.append(ir.Constant(member_type, ir.Undefined))
except Exception: # 回退:常量创建失败时使用零值
member_values.append(ir.Constant(member_type, 0))
if call_node.args:
for i, arg in enumerate(call_node.args):
if i < len(members):
_, member_type = members[i]
const = self._BuildScalarConstant(arg, member_type)
if const is not None:
member_values[i] = const
try:
return ir.Constant(struct_type, member_values)
except Exception: # 回退:结构体常量创建失败
return None