from __future__ import annotations import ast import re from typing import List import llvmlite.ir as ir from lib.includes import t from lib.includes.t import CTypeRegistry from lib.core.SymbolUtils import IsListAnnotation, AnnotationContainsName class DeclarationGenerator: """从 .pyi AST 生成纯 LLVM IR 声明(不使用字符串操作)""" # 匹配 LLVM 数组类型:[N x elem_type] _ARRAY_TYPE_RE: re.Pattern = re.compile(r'^\[(\d+)\s+x\s+(.+)\]$') @staticmethod def _decay_array_to_ptr(type_str: str) -> str: """C 语言中数组参数退化为指针:[N x elem_type] → elem_type*""" m = DeclarationGenerator._ARRAY_TYPE_RE.match(type_str) if m: return f'{m.group(2)}*' return type_str def __init__(self, struct_names: set[str] | None = None, enum_names: set[str] | None = None, module_sha1: str | None = None, target_triple: str | None = None, target_datalayout: str | None = None, struct_sha1_map: dict[str, str] | None = None, exception_names: set[str] | None = None, typedef_map: dict[str, ast.AST] | None = None, class_def_map: dict[str, ast.ClassDef] | None = None) -> None: self.ir: ir.Module = ir self.module: ir.Module | None = None self.builder: ir.IRBuilder | None = None self._DefineConstants: dict[str, int | str] = {} self.struct_names: set[str] = struct_names or set() self.enum_names: set[str] = enum_names or set() self.module_sha1: str | None = module_sha1 self.struct_sha1_map: dict[str, str] = struct_sha1_map or {} self.exception_names: set[str] = exception_names or set() self.target_triple: str = target_triple or "x86_64-none-elf" self.target_datalayout: str = target_datalayout or "e-m:e-p270:32:32-p271:32:32-p272:64:64-i64:64-f80:128-n8:16:32:64-S128" self.typedef_map: dict[str, ast.AST] = typedef_map or {} self._pyi_tree: ast.Module | None = None self._t_alias_map: dict[str, type] = {} # 跨模块类定义映射:{class_name: ClassDef_node},供 _get_inherited_members # 查找其他模块定义的父类(如 exprs.py 的 Name(AST) 查找 base.py 的 AST) self._cross_module_class_defs: dict[str, ast.ClassDef] = class_def_map or {} t.configure_platform(self.target_triple) def generate(self, pyi_content: str, src_path: str) -> str: tree: ast.Module = ast.parse(pyi_content) self._DefineConstants: dict[str, int | str] = {} self._pyi_tree = tree global_typedef_map: dict[str, ast.AST] = self.typedef_map self.typedef_map: dict[str, ast.AST] = {} if global_typedef_map: self.typedef_map.update(global_typedef_map) # 构建 t/c 类型别名映射(from t import CInt as myint) self._t_alias_map: dict[str, type] = {} for node in ast.iter_child_nodes(tree): if isinstance(node, ast.ImportFrom): if node.module in ('t', 'c'): for alias in node.names: name: str = alias.name asname: str = alias.asname or name if node.module == 't': t_cls: type | None = getattr(t, name, None) if isinstance(t_cls, type) and issubclass(t_cls, t.CType): self._t_alias_map[asname] = t_cls for node in ast.iter_child_nodes(tree): if isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name): var_name: str = node.target.id if node.value and isinstance(node.value, ast.Constant): self._DefineConstants[var_name] = node.value.value elif node.value and isinstance(node.value, ast.Name): self._DefineConstants[var_name] = node.value.id is_typedef: bool = False if isinstance(node.annotation, ast.Attribute) and hasattr(node.annotation, 'attr') and node.annotation.attr == 'CTypedef': is_typedef = True elif isinstance(node.annotation, ast.Name) and node.annotation.id == 'CTypedef': is_typedef = True elif isinstance(node.annotation, ast.BinOp) and isinstance(node.annotation.left, ast.Attribute) and hasattr(node.annotation.left, 'attr') and node.annotation.left.attr == 'CTypedef': is_typedef = True if is_typedef: if node.value: self.typedef_map[var_name] = node.value elif isinstance(node.annotation, ast.BinOp) and node.annotation.right: self.typedef_map[var_name] = node.annotation.right lines: list[str] = [] lines.append('; ModuleID = "transpyc_decl"') lines.append(f'target triple = "{self.target_triple}"') lines.append(f'target datalayout = "{self.target_datalayout}"') lines.append('') for node in ast.iter_child_nodes(tree): if isinstance(node, ast.FunctionDef): if hasattr(node, 'type_params') and node.type_params: continue decl: str | None = self._generate_func_decl(node) if decl: lines.append(decl) elif isinstance(node, ast.AnnAssign): decl: str | None = self._generate_global_decl(node) if decl: lines.append(decl) elif isinstance(node, ast.Assign): decl: str | None = self._generate_global_assign_decl(node) if decl: lines.append(decl) elif isinstance(node, ast.ClassDef): if hasattr(node, 'type_params') and node.type_params: continue decls: list[str] = self._generate_class_decl(node) lines.extend(decls) # 后处理:扫描引用但未定义的结构体类型,自动添加 opaque 声明 # 这解决了泛型类(如 list[T])在 stub 中被跳过导致的 "undefined type" 错误 opaque_decls: list[str] = self._collect_undefined_type_decls(lines) if opaque_decls: # 在 header 之后插入 opaque 声明(ModuleID, triple, datalayout, 空行 = 4 行) insert_idx: int = 4 for i, od in enumerate(opaque_decls): lines.insert(insert_idx + i, od) return '\n'.join(lines) def _collect_undefined_type_decls(self, lines: list[str]) -> list[str]: """扫描所有引用的 %"sha1.TypeName" 类型,返回未定义类型的 opaque 声明列表""" # 收集已定义的类型名 defined_types: set[str] = set() type_def_re: re.Pattern = re.compile(r'^(%".+?")\s*=\s*type\s') for line in lines: m = type_def_re.match(line.strip()) if m: defined_types.add(m.group(1)) # 扫描所有引用的类型名 referenced_types: set[str] = set() type_ref_re: re.Pattern = re.compile(r'%"[a-f0-9]+\.[^"]+"') for line in lines: for m in type_ref_re.finditer(line): ref: str = m.group(0) if ref not in defined_types: referenced_types.add(ref) # 生成 opaque 声明(排序确保确定性) result: list[str] = [] for type_name in sorted(referenced_types): result.append(f'{type_name} = type opaque') return result def _generate_func_decl(self, node: ast.FunctionDef) -> str | None: """生成函数声明""" func_name: str = node.name is_export: bool = self._is_export_func(node) if self.module_sha1 and not is_export: func_name = f"{self.module_sha1}.{func_name}" CReturnTypes: list[ast.AST] = [] if node.decorator_list: for decorator in node.decorator_list: if isinstance(decorator, ast.Call) and isinstance(decorator.func, ast.Attribute): if decorator.func.attr == 'CReturn': for arg in decorator.args: CReturnTypes.append(arg) if node.returns: if isinstance(node.returns, ast.Subscript) and isinstance(node.returns.value, ast.Name) and node.returns.value.id == 'tuple': slice_node: ast.AST = node.returns.slice if isinstance(slice_node, ast.Tuple): for elt in slice_node.elts: CReturnTypes.append(elt) else: CReturnTypes.append(slice_node) ret_type: str if CReturnTypes: elem_types: list[str] = [self._get_type_str(rt, embedded=True) for rt in CReturnTypes] ret_type = '{ ' + ', '.join(elem_types) + ' }' else: ret_type = self._get_type_str(node.returns) if not ret_type or ret_type == 'void': ret_type = 'void' elif ret_type == 'i8*': if node.returns is None: ret_type = 'void' params: list[str] = [] for arg in node.args.args: arg_type: str if arg.annotation: arg_type = self._get_type_str(arg.annotation) else: arg_type = 'i8*' # C 语言中数组参数退化为指针:[N x elem_type] → elem_type* arg_type = self._decay_array_to_ptr(arg_type) if arg_type and arg_type != 'void': params.append(arg_type) param_str: str = ', '.join(params) if params else '' if node.args.vararg: param_str = param_str + ', ...' if param_str else '...' if func_name[0].isdigit(): return f'declare {ret_type} @"{func_name}"({param_str})' return f'declare {ret_type} @{func_name}({param_str})' def _is_export_func(self, node: ast.FunctionDef) -> bool: """检查函数是否标记为 CExport""" if not node.returns: return False return self._check_annotation_for_export(node.returns) def _check_annotation_for_export(self, annotation: ast.AST | None) -> bool: """递归检查类型注解中是否包含 CExport 或 t.State""" return bool(annotation) and (AnnotationContainsName(annotation, 'CExport') or AnnotationContainsName(annotation, 'State')) def _generate_global_decl(self, node: ast.AnnAssign) -> str | None: """生成全局变量声明""" if not isinstance(node.target, ast.Name): return None var_name: str = node.target.id if isinstance(node.annotation, ast.Attribute) and hasattr(node.annotation, 'attr') and node.annotation.attr == 'CDefine': return None if isinstance(node.annotation, ast.Attribute) and hasattr(node.annotation, 'attr') and node.annotation.attr == 'CTypedef': return None if isinstance(node.annotation, ast.Name) and node.annotation.id == 'CTypedef': return None if isinstance(node.annotation, ast.BinOp) and isinstance(node.annotation.left, ast.Attribute) and hasattr(node.annotation.left, 'attr') and node.annotation.left.attr == 'CTypedef': return None if isinstance(node.annotation, ast.List): for elt in node.annotation.elts: if isinstance(elt, ast.Attribute) and hasattr(elt, 'attr') and elt.attr in ('CTypedef', 'Callable'): return None if isinstance(elt, ast.Subscript) and isinstance(elt.value, ast.Attribute) and isinstance(elt.value.value, ast.Name) and elt.value.value.id == 't' and elt.value.attr == 'Callable': return None VarType: str = self._get_type_str(node.annotation) if VarType.startswith('[0 x ') and node.value and isinstance(node.value, ast.List): elem_type: str = VarType[5:-1] actual_count: int = len(node.value.elts) if actual_count > 0: VarType = f'[{actual_count} x {elem_type}]' return f'@{var_name} = external global {VarType}' def _generate_global_assign_decl(self, node: ast.Assign) -> str | None: """从无类型标注赋值推断并生成全局变量声明""" if not node.targets or not isinstance(node.targets[0], ast.Name): return None var_name: str = node.targets[0].id VarType: str | None = self._infer_type(node.value) if VarType: return f'@{var_name} = external global {VarType}' return None def _resolve_base_kind(self, base_name: str) -> str | None: """解析基类名称的类型种类(支持别名) Returns: 'enum' | 'union' | 'renum' | 'struct' | None """ # 1) 查 t 别名映射(from t import CEnum as en) if base_name in self._t_alias_map: t_cls: type = self._t_alias_map[base_name] if issubclass(t_cls, (t.CEnum, t.REnum)): return 'renum' if issubclass(t_cls, t.REnum) else 'enum' if issubclass(t_cls, t.CUnion): return 'union' if issubclass(t_cls, t.CStruct): return 'struct' # 2) 直接查 t 模块属性 t_cls: type | None = getattr(t, base_name, None) if isinstance(t_cls, type) and issubclass(t_cls, t.CType): if issubclass(t_cls, (t.CEnum, t.REnum)): return 'renum' if issubclass(t_cls, t.REnum) else 'enum' if issubclass(t_cls, t.CUnion): return 'union' if issubclass(t_cls, t.CStruct): return 'struct' return None def _is_marker_base(self, base_name: str) -> bool: """检查是否为标记基类(Object/CVTable/Exception/CEnum/CUnion/CStruct/REnum 及其别名)""" if base_name in ('Object', 'CVTable', 'Exception', 'CEnum', 'Enum', 'CStruct', 'CUnion', 'REnum'): return True if self._resolve_base_kind(base_name) is not None: return True return False def _generate_class_decl(self, node: ast.ClassDef) -> List[str]: """生成结构体/类声明""" decls: list[str] = [] class_name: str = node.name is_enum: bool = False is_renum: bool = False is_union: bool = False is_exception: bool = False if node.bases: for base in node.bases: base_name: str | None = None if isinstance(base, ast.Attribute) and hasattr(base, 'attr'): base_name = base.attr elif isinstance(base, ast.Name) and hasattr(base, 'id'): base_name = base.id if base_name: if base_name in self.exception_names or base_name == 'Exception': is_exception = True break kind: str | None = self._resolve_base_kind(base_name) if kind == 'enum': is_enum = True break elif kind == 'union': is_union = True elif kind == 'renum': is_renum = True if is_exception: return decls if is_enum: enum_values: dict[str, int] = {} next_value: int = 0 for item in node.body: if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name): var_name: str = item.target.id if item.value and isinstance(item.value, ast.Constant) and isinstance(item.value.value, int): enum_values[var_name] = item.value.value else: enum_values[var_name] = next_value next_value += 1 elif isinstance(item, ast.Assign) and len(item.targets) == 1 and isinstance(item.targets[0], ast.Name): var_name: str = item.targets[0].id if item.value and isinstance(item.value, ast.Constant) and isinstance(item.value.value, int): enum_values[var_name] = item.value.value next_value = item.value.value + 1 else: enum_values[var_name] = next_value next_value += 1 for var_name, value in enum_values.items(): decls.append(f'@__config_{class_name}_{var_name} = external global i32') return decls if is_renum: # REnum 布局: { i32 __tag, } # 对齐 _EmitREnumLlvm (HandlesClassDef.py) 的 Phase 2 逻辑 # Bug 修复:按每个位置的 max 字段尺寸构建布局,而非选总尺寸最大的变体。 # 不同变体在同一位置可能有不同尺寸的字段(如 Concrete 位置3是 i32, # 但 Var 位置3是 MetaType*),用最大变体的布局会导致指针截断。 variant_fields_list: list[list[str]] = [] for item in node.body: if isinstance(item, ast.ClassDef): variant_fields: list[str] = [] for nested_item in item.body: if isinstance(nested_item, ast.AnnAssign) and isinstance(nested_item.target, ast.Name): ft: str = self._get_type_str(nested_item.annotation, embedded=True) if ft: variant_fields.append(ft) variant_fields_list.append(variant_fields) elif isinstance(item, ast.Assign): # 显式值变体(无 payload) variant_fields_list.append([]) # 计算每个位置的最大尺寸 max_field_count: int = 0 for vf in variant_fields_list: if len(vf) > max_field_count: max_field_count = len(vf) renum_member_types: list[str] = ['i32'] # __tag for pos in range(max_field_count): pos_max_size: int = 0 for vf in variant_fields_list: if pos < len(vf): sz: int = self._llvm_type_size_str(vf[pos]) if sz > pos_max_size: pos_max_size = sz # <=4 用 i32,>4 用 i64(8字节,可容纳指针) if pos_max_size <= 4: renum_member_types.append('i32') else: renum_member_types.append('i64') if len(renum_member_types) < 2: # 无类变体或所有变体无字段时,使用 i32 作为占位(对齐 _EmitREnumLlvm 默认行为) renum_member_types = ['i32', 'i32'] struct_type_name: str = f'%"{self.module_sha1}.{class_name}"' if self.module_sha1 else f'%struct.{class_name}' struct_decl: str = f'{struct_type_name} = type {{ {", ".join(renum_member_types)} }}' decls.append(struct_decl) return decls member_types: list[str] = [] has_bitfield: bool = False total_bits: int = 0 has_methods: bool = any(isinstance(item, ast.FunctionDef) for item in node.body) is_cvtable: bool = False is_cpython_object: bool = False is_novtable: bool = False if hasattr(node, 'decorator_list') and node.decorator_list: for decorator in node.decorator_list: if isinstance(decorator, ast.Attribute): if getattr(decorator.value, 'id', None) == 't': if decorator.attr == 'CVTable': is_cvtable = True elif decorator.attr == 'Object': is_cpython_object = True elif decorator.attr == 'NoVTable': is_novtable = True elif isinstance(decorator, ast.Name): if decorator.id == 'CVTable': is_cvtable = True elif decorator.id == 'Object': is_cpython_object = True elif decorator.id == 'NoVTable': is_novtable = True if has_methods and not is_cpython_object: is_cpython_object = True has_parent_class: bool = False parent_is_novtable: bool = False for base in node.bases: base_name: str | None = None if isinstance(base, ast.Attribute): base_name = base.attr elif isinstance(base, ast.Name): base_name = base.id elif isinstance(base, ast.Subscript): if isinstance(base.value, ast.Attribute): base_name = base.value.attr elif isinstance(base.value, ast.Name): base_name = base.value.id if base_name and not self._is_marker_base(base_name): has_parent_class = True # 检查父类是否为 @t.NoVTable(非多态继承) _p_node: ast.ClassDef | None = None for n in ast.iter_child_nodes(self._pyi_tree): if isinstance(n, ast.ClassDef) and n.name == base_name: _p_node = n break if _p_node is None: _p_node = self._cross_module_class_defs.get(base_name) if _p_node is not None and hasattr(_p_node, 'decorator_list') and _p_node.decorator_list: for dec in _p_node.decorator_list: if isinstance(dec, ast.Attribute) and getattr(dec.value, 'id', None) == 't' and dec.attr == 'NoVTable': parent_is_novtable = True elif isinstance(dec, ast.Name) and dec.id == 'NoVTable': parent_is_novtable = True break # 仅当非 NoVTable 且父类非 NoVTable 时才因继承启用 CVTable # (@t.NoVTable 标记的类及其子类不使用 vtable,字段展平嵌入) if has_parent_class and not is_cvtable and not is_novtable and not parent_is_novtable: is_cvtable = True base_has_vtable: bool = False for base in node.bases: base_name: str | None = None if isinstance(base, ast.Attribute): base_name = base.attr elif isinstance(base, ast.Name): base_name = base.id elif isinstance(base, ast.Subscript): if isinstance(base.value, ast.Attribute): base_name = base.value.attr elif isinstance(base.value, ast.Name): base_name = base.value.id if base_name and not self._is_marker_base(base_name): # 先在当前 .pyi 中查找父类,找不到则查跨模块类定义映射 _base_node: ast.ClassDef | None = None for n in ast.iter_child_nodes(self._pyi_tree): if isinstance(n, ast.ClassDef) and n.name == base_name: _base_node = n break if _base_node is None: _base_node = self._cross_module_class_defs.get(base_name) if _base_node is not None and hasattr(_base_node, 'decorator_list') and _base_node.decorator_list: for dec in _base_node.decorator_list: if isinstance(dec, ast.Attribute) and getattr(dec.value, 'id', None) == 't' and dec.attr == 'CVTable': base_has_vtable = True elif isinstance(dec, ast.Name) and dec.id == 'CVTable': base_has_vtable = True if base_has_vtable: break # NoVTable 类不添加 vtable 指针(即使有方法和父类) if has_methods and is_cvtable and not base_has_vtable and not is_novtable: member_types.append('i8*') seen_member_names: set[str] = set() for base in node.bases: base_name: str | None = None type_args: list[str] = [] if isinstance(base, ast.Attribute): base_name = base.attr elif isinstance(base, ast.Name): base_name = base.id elif isinstance(base, ast.Subscript): if isinstance(base.value, ast.Attribute): base_name = base.value.attr elif isinstance(base.value, ast.Name): base_name = base.value.id slice_node: ast.AST = base.slice if isinstance(slice_node, ast.Name): type_args.append(slice_node.id) elif isinstance(slice_node, ast.Tuple): for elt in slice_node.elts: if isinstance(elt, ast.Name): type_args.append(elt.id) if base_name and not self._is_marker_base(base_name): base_member_types: list[str] base_seen: set[str] base_member_types, base_seen = self._get_inherited_members(base_name) if type_args: base_member_types = self._specialize_member_types(base_member_types, type_args) for mt in base_member_types: member_types.append(mt) seen_member_names.update(base_seen) for item in node.body: if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name): seen_member_names.add(item.target.id) init_members: list[ast.AnnAssign] = [] for item in node.body: if isinstance(item, ast.FunctionDef) and item.name == '__init__': for stmt in item.body: if (isinstance(stmt, ast.AnnAssign) and isinstance(stmt.target, ast.Attribute) and isinstance(stmt.target.value, ast.Name) and stmt.target.value.id == 'self'): init_members.append(stmt) untyped_self_members: list[ast.Assign] = [] for m_name in [m.attr for m in init_members if isinstance(m.target, ast.Attribute)]: seen_member_names.add(m_name) for item in node.body: if isinstance(item, ast.FunctionDef): if hasattr(item, 'type_params') and item.type_params: continue for stmt in item.body: if (isinstance(stmt, ast.Assign) and len(stmt.targets) == 1 and isinstance(stmt.targets[0], ast.Attribute) and isinstance(stmt.targets[0].value, ast.Name) and stmt.targets[0].value.id == 'self'): attr_name: str = stmt.targets[0].attr if attr_name not in seen_member_names: seen_member_names.add(attr_name) untyped_self_members.append(stmt) for item in list(node.body) + init_members + untyped_self_members: if isinstance(item, ast.AnnAssign): # 跳过编译期元数据字段(__provides__/__requires__/__require_must__) if isinstance(item.target, ast.Name) and item.target.id in ('__provides__', '__requires__', '__require_must__'): continue bit_width: int | None = self._get_bitfield_width(item.annotation) if bit_width is not None: has_bitfield = True total_bits += bit_width else: VarType: str = self._get_type_str(item.annotation, embedded=True) if IsListAnnotation(item.annotation): slice_node: ast.AST = item.annotation.slice is_dynamic: bool = False if isinstance(slice_node, ast.Tuple) and len(slice_node.elts) == 2: count_node: ast.AST = slice_node.elts[1] if isinstance(count_node, ast.Constant) and count_node.value is None: is_dynamic = True elif not (isinstance(count_node, ast.Constant) and isinstance(count_node.value, int) and count_node.value > 0): if not isinstance(count_node, (ast.Name, ast.BinOp)): is_dynamic = True elif not isinstance(slice_node, ast.Tuple): is_dynamic = True if is_dynamic: elem_type: str = self._get_type_str(slice_node if not isinstance(slice_node, ast.Tuple) else slice_node.elts[0], embedded=True) init_len: int = 0 if isinstance(item.value, ast.List): init_len = len(item.value.elts) elif isinstance(item.value, ast.Constant) and isinstance(item.value.value, str): init_len = len(item.value.value) + 1 VarType = f'[{init_len} x {elem_type}]' else: # 固定大小数组: t.CArray[elem_type, count] → [count x elem_type] if isinstance(slice_node, ast.Tuple) and len(slice_node.elts) == 2: elem_type: str = self._get_type_str(slice_node.elts[0], embedded=True) count_val: int = self._get_const_int(slice_node.elts[1]) if count_val > 0: VarType = f'[{count_val} x {elem_type}]' elif isinstance(slice_node, (ast.Attribute, ast.Name, ast.Subscript)): # 单参数 t.CArray[elem_type] → 指针模式 elem_type: str = self._get_type_str(slice_node, embedded=True) VarType = f'{elem_type}*' member_types.append(VarType) elif isinstance(item, ast.Assign): if (len(item.targets) == 1 and isinstance(item.targets[0], ast.Attribute) and isinstance(item.targets[0].value, ast.Name) and item.targets[0].value.id == 'self'): attr_name: str = item.targets[0].attr if attr_name in seen_member_names: continue InferredType: str = 'i32' if item.value and isinstance(item.value, ast.Constant): if isinstance(item.value.value, float): InferredType = 'float' elif isinstance(item.value.value, bool): InferredType = 'i8' elif isinstance(item.value.value, str): InferredType = 'i8*' member_types.append(InferredType) else: member_types.append('i32') if has_bitfield: if total_bits <= 8: member_types.insert(0, 'i8') elif total_bits <= 16: member_types.insert(0, 'i16') elif total_bits <= 32: member_types.insert(0, 'i32') else: member_types.insert(0, 'i64') struct_type_name: str = f'%"{self.module_sha1}.{class_name}"' if self.module_sha1 else f'%struct.{class_name}' struct_decl: str = f'{struct_type_name} = type {{ {", ".join(member_types)} }}' decls.append(struct_decl) if is_cpython_object or is_cvtable: new_func_name: str = f'{class_name}.__before_init__' if self.module_sha1: new_func_name = f"{self.module_sha1}.{new_func_name}" new_func_decl: str if new_func_name[0].isdigit(): new_func_decl = f'declare void @"{new_func_name}"({struct_type_name}*)' else: new_func_decl = f'declare void @{new_func_name}({struct_type_name}*)' decls.append(new_func_decl) for item in node.body: if isinstance(item, ast.FunctionDef): if hasattr(item, 'type_params') and item.type_params: continue method_name: str = f'{class_name}.{item.name}' if self.module_sha1: method_name = f"{self.module_sha1}.{method_name}" ret_type: str = self._get_type_str(item.returns) if item.returns else 'void' if not ret_type: ret_type = 'void' params: list[str] = [] for arg_idx, arg in enumerate(item.args.args): arg_type: str if arg.annotation: arg_type = self._get_type_str(arg.annotation) elif arg_idx == 0 and arg.arg == 'self': arg_type = f'{struct_type_name}*' else: arg_type = 'i8*' # C 语言中数组参数退化为指针:[N x elem_type] → elem_type* arg_type = self._decay_array_to_ptr(arg_type) params.append(arg_type) param_str: str = ', '.join(params) if params else '' if method_name[0].isdigit(): decls.append(f'declare {ret_type} @"{method_name}"({param_str})') else: decls.append(f'declare {ret_type} @{method_name}({param_str})') # 为继承但未覆写的方法生成包装声明 # 这些声明让跨模块调用能正确解析子类包装函数的签名, # 避免 stub 缺失导致默认 i32 返回类型 → 64 位指针截断 decls.extend(self._generate_inherited_method_decls(node, class_name, struct_type_name)) return decls def _generate_inherited_method_decls(self, node: ast.ClassDef, child_class_name: str, child_struct_type: str) -> list[str]: """为继承但未覆写的方法生成包装声明 在 Phase 1 stub 中声明子类继承方法的签名(与父类相同,但 self 参数为子类类型)。 这样测试模块翻译时能从 temp stub 中获取正确的返回类型, 避免 Phase 2 includes 翻译滞后导致跨模块调用使用默认 i32 返回类型。 """ decls: list[str] = [] # 收集子类自身定义的方法名(这些不需要生成包装,子类有自己的实现) seen_method_names: set[str] = set() for item in node.body: if isinstance(item, ast.FunctionDef): seen_method_names.add(item.name) # 沿继承链收集父类方法,近祖先优先 for base in node.bases: base_name: str | None = None if isinstance(base, ast.Attribute): base_name = base.attr elif isinstance(base, ast.Name): base_name = base.id elif isinstance(base, ast.Subscript): if isinstance(base.value, ast.Attribute): base_name = base.value.attr elif isinstance(base.value, ast.Name): base_name = base.value.id if not base_name or self._is_marker_base(base_name): continue self._CollectInheritedWrapperDecls( base_name, child_class_name, child_struct_type, seen_method_names, decls) return decls def _CollectInheritedWrapperDecls( self, parent_name: str, child_class_name: str, child_struct_type: str, seen_method_names: set[str], decls: list[str]) -> None: """递归收集父类及其祖先的方法,为未覆写的方法生成包装声明""" # 查找父类节点:先在当前 .pyi 中查找,找不到则查跨模块类定义映射 parent_node: ast.ClassDef | None = None for n in ast.iter_child_nodes(self._pyi_tree): if isinstance(n, ast.ClassDef) and n.name == parent_name: parent_node = n break if parent_node is None: parent_node = self._cross_module_class_defs.get(parent_name) if parent_node is None: return # 为父类的每个方法生成包装声明 for item in parent_node.body: if isinstance(item, ast.FunctionDef): if hasattr(item, 'type_params') and item.type_params: continue method_name: str = item.name # __before_init__ 只在自身类生成,不生成继承包装 if method_name == '__before_init__': continue # 跳过子类已覆写的方法,以及已处理的祖先方法 if method_name in seen_method_names: continue seen_method_names.add(method_name) # 构建包装方法名:{module_sha1}.{child_class}.{method_name} wrapper_name: str = f'{child_class_name}.{method_name}' if self.module_sha1: wrapper_name = f"{self.module_sha1}.{wrapper_name}" # 解析返回类型(与父类方法相同) ret_type: str = self._get_type_str(item.returns) if item.returns else 'void' if not ret_type: ret_type = 'void' elif ret_type == 'i8*': if item.returns is None: ret_type = 'void' # 构建参数列表:第一个参数(self)替换为子类类型,其余与父类相同 params: list[str] = [] for arg_idx, arg in enumerate(item.args.args): arg_type: str if arg_idx == 0 and arg.arg == 'self': arg_type = f'{child_struct_type}*' else: if arg.annotation: arg_type = self._get_type_str(arg.annotation) else: arg_type = 'i8*' arg_type = self._decay_array_to_ptr(arg_type) params.append(arg_type) param_str: str = ', '.join(params) if params else '' if wrapper_name[0].isdigit(): decls.append(f'declare {ret_type} @"{wrapper_name}"({param_str})') else: decls.append(f'declare {ret_type} @{wrapper_name}({param_str})') # 递归处理祖先类 for base in parent_node.bases: base_name: str | None = None if isinstance(base, ast.Attribute): base_name = base.attr elif isinstance(base, ast.Name): base_name = base.id elif isinstance(base, ast.Subscript): if isinstance(base.value, ast.Attribute): base_name = base.value.attr elif isinstance(base.value, ast.Name): base_name = base.value.id if base_name and not self._is_marker_base(base_name): self._CollectInheritedWrapperDecls( base_name, child_class_name, child_struct_type, seen_method_names, decls) def _get_inherited_members(self, base_name: str) -> tuple[list[str], set[str]]: member_types: list[str] = [] seen_names: set[str] = set() if not hasattr(self, '_pyi_tree') or self._pyi_tree is None: return member_types, seen_names base_node: ast.ClassDef | None = None for node in ast.iter_child_nodes(self._pyi_tree): if isinstance(node, ast.ClassDef) and node.name == base_name: base_node = node break if base_node is None: # 跨模块父类查找:当前 .pyi 中没有父类定义时,从跨模块类定义映射中查找 # (如 exprs.py 的 Name(AST),AST 定义在 base.py) base_node = self._cross_module_class_defs.get(base_name) if base_node is None: return member_types, seen_names base_decls: list[str] = self._generate_class_decl(base_node) for decl in base_decls: if '= type {' in decl: type_body: str = decl.split('= type {')[1].rstrip('}').strip() if type_body: for t in type_body.split(','): t = t.strip() if t: member_types.append(t) break for item in base_node.body: if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name): seen_names.add(item.target.id) elif isinstance(item, ast.Assign) and len(item.targets) == 1: if isinstance(item.targets[0], ast.Attribute) and isinstance(item.targets[0].value, ast.Name) and item.targets[0].value.id == 'self': seen_names.add(item.targets[0].attr) for item in base_node.body: if isinstance(item, ast.FunctionDef) and item.name == '__init__': for stmt in item.body: if isinstance(stmt, ast.AnnAssign) and isinstance(stmt.target, ast.Attribute) and isinstance(stmt.target.value, ast.Name) and stmt.target.value.id == 'self': seen_names.add(stmt.target.attr) elif isinstance(stmt, ast.Assign) and len(stmt.targets) == 1 and isinstance(stmt.targets[0], ast.Attribute) and isinstance(stmt.targets[0].value, ast.Name) and stmt.targets[0].value.id == 'self': seen_names.add(stmt.targets[0].attr) return member_types, seen_names def _specialize_member_types(self, member_types: list[str], type_args: list[str]) -> list[str]: """将基类成员类型中的泛型参数 T 替换为具体类型参数 处理泛型基类继承(如 GSListNode[BasicBlock])时,基类 _get_inherited_members 返回的成员类型中会包含未特化的 T 引用(形式为 %"SHA1.T"*)。 替换策略:使用 i8*(不透明指针)代替 T*,而非具体类型指针。 原因:跨模块场景中,具体类型指针的 SHA1 前缀可能与导入模块不一致, 导致 TransPyC 类型转换失败。i8* 作为通用指针类型,可安全赋值给 任何 Class | t.CPtr 注解的变量,运行时行为与 T* 相同(都是指针)。 Args: member_types: 基类成员类型字符串列表 type_args: 类型参数名列表(如 ['BasicBlock']) Returns: 特化后的成员类型字符串列表 """ if not type_args: return member_types result: list[str] = [] for mt in member_types: new_mt: str = mt # 将 %"SHA1.T"* 替换为 i8*(不透明指针) new_mt = re.sub(r'%"[a-f0-9]+\.T"\*', 'i8*', new_mt) # 将 %"SHA1.T"(非指针,如 embedded struct)替换为 i8* new_mt = re.sub(r'%"[a-f0-9]+\.T"', 'i8*', new_mt) result.append(new_mt) return result def _llvm_type_size_str(self, type_str: str) -> int: """估算 LLVM 类型字符串的大小(字节),用于 REnum 变体大小比较 对齐 _EmitREnumLlvm 中的大小计算逻辑: - i1=1, i8=1, i16=2, i32=4, i64=8, float=4, double=8, ptr=8 - [N x T] = N * size(T) - 其他结构体类型按指针大小 8 估算(保守值) """ s: str = type_str.strip() if not s: return 0 # 指针类型 if s.endswith('*'): return 8 # 数组类型 [N x T] m = self._ARR_TYPE_RE.match(s) if hasattr(self, '_ARR_TYPE_RE') else None if m: try: n: int = int(m.group(1)) return n * self._llvm_type_size_str(m.group(2)) except ValueError: return 8 # 基本整数/浮点类型 size_map: dict[str, int] = { 'i1': 1, 'i8': 1, 'i16': 2, 'i32': 4, 'i64': 8, 'float': 4, 'double': 8, 'half': 2, 'fp128': 16, 'void': 0, } if s in size_map: return size_map[s] # %{...} 结构体类型 - 保守估算为指针大小 if s.startswith('%') or s.startswith('{'): return 8 # 其他未知类型保守估算 return 8 def _get_type_str(self, annotation: ast.AST | None, embedded: bool = False) -> str: if annotation is None: return 'i8*' non_struct_types: set[str] = {'list', 'dict', 'tuple', 'set', 'array', 'CArray'} if isinstance(annotation, ast.Name): if annotation.id == 'None': return 'void' if annotation.id in ('str', 'bytes'): return 'i8*' llvm_type: str | None = CTypeRegistry.NameToLLVM(annotation.id) if llvm_type is not None: return llvm_type resolved: tuple[type, int] | None = CTypeRegistry.ResolveName(annotation.id) if resolved is not None: ctype_cls: type ptr_level: int ctype_cls, ptr_level = resolved base: str = CTypeRegistry.CTypeToLLVM(ctype_cls) if ptr_level > 0: if base == 'void': return 'i8*' if '*' in base: return base return f'{base}*' return base if annotation.id in self.struct_names: sha1: str | None = self.struct_sha1_map.get(annotation.id, self.module_sha1) sname: str = f'%"{sha1}.{annotation.id}"' if sha1 else f'%struct.{annotation.id}' if embedded: return sname else: return f'{sname}*' if annotation.id in non_struct_types: return 'i32' if annotation.id in self.enum_names: return 'i32' if annotation.id in self.typedef_map: resolved_str: str = self._get_type_str(self.typedef_map[annotation.id], embedded=embedded) return resolved_str # 检查 t 类型别名(from t import CInt as myint) if annotation.id in self._t_alias_map: t_cls: type = self._t_alias_map[annotation.id] llvm: str | None = CTypeRegistry.CTypeToLLVM(t_cls) if llvm: return llvm sha1: str | None = self.struct_sha1_map.get(annotation.id, self.module_sha1) sname: str = f'%"{sha1}.{annotation.id}"' if sha1 else f'%struct.{annotation.id}' return f'{sname}*' elif isinstance(annotation, ast.Attribute): attr_name: str = annotation.attr if hasattr(annotation, 'attr') else '' module_name: str = '' if hasattr(annotation, 'value'): if isinstance(annotation.value, ast.Name): module_name = annotation.value.id elif isinstance(annotation.value, ast.Attribute) and hasattr(annotation.value, 'attr'): module_name = annotation.value.attr if (module_name == 't' and attr_name == 'State') or attr_name == 'State': return '' if (module_name == 't' and attr_name == 'Callable') or attr_name == 'Callable': return 'i8*' llvm_type: str | None = CTypeRegistry.NameToLLVM(attr_name) if llvm_type is not None: return llvm_type resolved: tuple[type, int] | None = CTypeRegistry.ResolveName(attr_name) if resolved is not None: ctype_cls: type ptr_level: int ctype_cls, ptr_level = resolved base: str = CTypeRegistry.CTypeToLLVM(ctype_cls) if ptr_level > 0: if base == 'void': return 'i8*' if '*' in base: return base return f'{base}*' return base if attr_name in self.typedef_map: return self._get_type_str(self.typedef_map[attr_name], embedded=embedded) if attr_name in self.enum_names: return 'i32' if attr_name in self.struct_names: sha1: str | None = self.struct_sha1_map.get(attr_name, self.module_sha1) sname: str = f'%"{sha1}.{attr_name}"' if sha1 else f'%struct.{attr_name}' if embedded: return sname else: return f'{sname}*' if attr_name and attr_name[0].isupper() and attr_name not in non_struct_types: sha1: str | None = self.struct_sha1_map.get(attr_name, self.module_sha1) sname: str = f'%"{sha1}.{attr_name}"' if sha1 else f'%struct.{attr_name}' return f'{sname}*' if attr_name and attr_name not in non_struct_types: sha1: str | None = self.struct_sha1_map.get(attr_name, self.module_sha1) sname: str = f'%"{sha1}.{attr_name}"' if sha1 else f'%struct.{attr_name}' return f'{sname}*' return 'i32' elif isinstance(annotation, ast.BinOp): left_type: str = self._get_type_str(annotation.left, embedded=embedded) right_type: str = self._get_type_str(annotation.right, embedded=embedded) if left_type in ('', 'void'): return right_type if right_type not in ('', 'void') else (left_type or right_type) if right_type in ('', 'void'): return left_type # t.CPtr (i8*) 与 left 合并:对齐 Phase 2 _MergeAllComponentTypes 的 # "first CPtr absorption" 语义:struct | t.CPtr → struct*(吸收,struct 变量默认是指针)。 # 但对于本身已是指针的类型(如 bytes=str=i8*),bytes | t.CPtr → i8**(不吸收)。 # 例外:泛型特化类型(含 '[')回退为 i8*,避免引用未定义的特化结构体。 if right_type == 'i8*': if '[' in left_type: return 'i8*' if left_type in ('', 'void'): return 'i8*' if left_type.startswith('%"') and left_type.endswith('*'): return left_type if '*' in left_type: return f'{left_type}*' return f'{left_type}*' if '*' in left_type: return left_type if '*' in right_type: return right_type return left_type elif isinstance(annotation, ast.Subscript): base: str = self._get_type_str(annotation.value) if base == 'i8*' and isinstance(annotation.slice, ast.Constant): return f'[{self._get_const_int(annotation.slice)} x i8]' # 处理 t.CPtr[Type] 语法(指向 Type 的指针) # 例如: t.CPtr[t.CPtr] -> i8** (void**), t.CPtr[t.CInt] -> i32* (int*) # t.CPtr[t.CPtr[t.CInt]] -> i32** (int**) ValueIsCPtr: bool = ( (isinstance(annotation.value, ast.Attribute) and annotation.value.attr == 'CPtr') or (isinstance(annotation.value, ast.Name) and annotation.value.id == 'CPtr') ) if ValueIsCPtr and isinstance(annotation.slice, (ast.Attribute, ast.Name, ast.Subscript)): slice_type: str = self._get_type_str(annotation.slice, embedded=True) if slice_type and slice_type != 'void': return f'{slice_type}*' return 'i8*' if isinstance(annotation.slice, ast.Constant) and isinstance(annotation.slice.value, int): return f'[{annotation.slice.value} x {base}]' # 处理 t.CArray[elem_type, count] 注解 → [count x elem_type] # 支持 t.CArray (ast.Attribute) 和 CArray (ast.Name) 两种写法 ValueIsCArray: bool = ( (isinstance(annotation.value, ast.Attribute) and annotation.value.attr == 'CArray') or (isinstance(annotation.value, ast.Name) and annotation.value.id == 'CArray') ) if ValueIsCArray: SliceNode: ast.AST = annotation.slice if isinstance(SliceNode, ast.Tuple) and len(SliceNode.elts) == 2: ElemType: str = self._get_type_str(SliceNode.elts[0], embedded=True) CountVal: int = self._get_const_int(SliceNode.elts[1]) if CountVal > 0: return f'[{CountVal} x {ElemType}]' # count 为 None 或 0 → 零长度数组(外部符号) return f'[0 x {ElemType}]' elif isinstance(SliceNode, (ast.Attribute, ast.Name, ast.Subscript)): # 单参数 t.CArray[elem_type] → 指针模式 ElemType: str = self._get_type_str(SliceNode, embedded=True) return f'{ElemType}*' if isinstance(annotation.value, ast.Name) and annotation.value.id == 'tuple': slice_node: ast.AST = annotation.slice elem_types: list[str] if isinstance(slice_node, ast.Tuple): elem_types = [self._get_type_str(e, embedded=True) for e in slice_node.elts] else: elem_types = [self._get_type_str(slice_node, embedded=True)] return '{ ' + ', '.join(elem_types) + ' }' # 处理泛型类类型注解(如 list[str]) if isinstance(annotation.value, ast.Name): gc_name: str = annotation.value.id if gc_name in non_struct_types and gc_name != 'tuple': gc_slice: ast.AST = annotation.slice gc_args: list[str] = [] if isinstance(gc_slice, ast.Name): gc_args.append(gc_slice.id) elif isinstance(gc_slice, ast.Tuple): gc_valid: bool = True for gc_elt in gc_slice.elts: if isinstance(gc_elt, ast.Name): gc_args.append(gc_elt.id) else: gc_valid = False break if not gc_valid: gc_args = [] if gc_args: gc_mangled: list[str] = [] for gc_ta in gc_args: gc_llvm: str | None = CTypeRegistry.NameToLLVM(gc_ta) if gc_llvm: if gc_llvm in ('float', 'double', 'half', 'fp128'): gc_cls: type | None = CTypeRegistry.GetClassByName(gc_ta) gc_sz: int = gc_cls().Size if gc_cls else 0 gc_mangled.append('f' + str(gc_sz)) else: gc_mangled.append(gc_llvm) else: gc_mangled.append(gc_ta) gc_spec: str = gc_name + '[' + ']['.join(gc_mangled) + ']' gc_sha1: str | None = self.module_sha1 gc_sname: str = f'%"{gc_sha1}.{gc_spec}"' if gc_sha1 else f'%struct.{gc_spec}' if embedded: return gc_sname return f'{gc_sname}*' return f'{base}' elif isinstance(annotation, ast.Constant): if annotation.value is None: return 'void' if isinstance(annotation.value, int): return 'i32' if isinstance(annotation.value, str): type_name: str = annotation.value if type_name in self.enum_names: return 'i32' if type_name in self.struct_names: sha1: str | None = self.struct_sha1_map.get(type_name, self.module_sha1) sname: str = f'%"{sha1}.{type_name}"' if sha1 else f'%struct.{type_name}' if embedded: return sname else: return f'{sname}*' llvm_type: str | None = CTypeRegistry.NameToLLVM(type_name) if llvm_type is not None: return llvm_type resolved: tuple[type, int] | None = CTypeRegistry.ResolveName(type_name) if resolved is not None: ctype_cls: type ptr_level: int ctype_cls, ptr_level = resolved base: str = CTypeRegistry.CTypeToLLVM(ctype_cls) if ptr_level > 0: if base == 'void': return 'i8*' if '*' in base: return base return f'{base}*' return base sha1: str | None = self.struct_sha1_map.get(type_name, self.module_sha1) sname: str = f'%"{sha1}.{type_name}"' if sha1 else f'%struct.{type_name}' return f'{sname}*' return 'i8*' elif isinstance(annotation, ast.Call): if isinstance(annotation.func, ast.Name) and annotation.func.id == 'callable': return 'i8*' if isinstance(annotation.func, ast.Attribute) and annotation.func.attr == 'Callable': return 'i8*' return 'i32' def _infer_type(self, value: ast.AST) -> str: """从值推断类型""" if isinstance(value, ast.Constant): if isinstance(value.value, int): return 'i32' elif isinstance(value.value, float): return 'double' elif isinstance(value.value, str): return 'i8*' elif isinstance(value.value, bool): return 'i8' elif isinstance(value, ast.List): return 'i8*' elif isinstance(value, ast.Dict): return 'i8*' elif isinstance(value, ast.Name): return 'i32' elif isinstance(value, ast.BinOp): return self._infer_type(value.left) elif isinstance(value, ast.Call): return 'i8*' return 'i32' def _get_bitfield_width(self, annotation: ast.AST) -> int | None: if isinstance(annotation, ast.BinOp) and isinstance(annotation.op, ast.BitOr): right_width: int | None = self._get_bitfield_width(annotation.right) if right_width is not None: return right_width return self._get_bitfield_width(annotation.left) if isinstance(annotation, ast.Call): if isinstance(annotation.func, ast.Attribute) and annotation.func.attr == 'Bit': if annotation.args: arg: ast.AST = annotation.args[0] if isinstance(arg, ast.Constant) and isinstance(arg.value, int): return arg.value if isinstance(annotation, ast.Subscript): if isinstance(annotation.value, ast.Attribute) and annotation.value.attr == 'Bit': if isinstance(annotation.slice, ast.Constant) and isinstance(annotation.slice.value, int): return annotation.slice.value return None def _get_const_int(self, node: ast.AST) -> int: """获取常量整数值,支持符号常量和简单表达式""" if isinstance(node, ast.Constant) and isinstance(node.value, int): return node.value if isinstance(node, ast.Name): if node.id in self._DefineConstants: val: int | str = self._DefineConstants[node.id] if isinstance(val, int): return val if isinstance(node, ast.BinOp): left_val: int = self._get_const_int(node.left) right_val: int = self._get_const_int(node.right) if left_val and right_val: if isinstance(node.op, ast.Add): return left_val + right_val if isinstance(node.op, ast.Sub): return left_val - right_val if isinstance(node.op, ast.Mult): return left_val * right_val if isinstance(node.op, ast.Div): return left_val // right_val if isinstance(node.op, ast.FloorDiv): return left_val // right_val return 0