966 lines
51 KiB
Python
966 lines
51 KiB
Python
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
|
||
from lib.core.SymbolNode import SymbolNode
|
||
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[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}"
|
||
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 = SymbolNode.CreateClass(
|
||
name=ClassName,
|
||
TypeKind='union',
|
||
lineno=Node.lineno,
|
||
file='<stdin>'
|
||
)
|
||
self.Trans.SymbolTable[ClassName] = UnionNode.attributes
|
||
|
||
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[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[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[ClassName] = EnumTypeInfo
|
||
from lib.core.export_table import EnumMember as ExportEnumMember
|
||
enum_export = self.Trans.ExportTable.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 value in self.Trans.SymbolTable:
|
||
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[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 value in self.Trans.SymbolTable:
|
||
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[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
|