1201 lines
63 KiB
Python
1201 lines
63 KiB
Python
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):
|
||
# PEP 695 泛型类不再完全跳过:子类(如 Value(GSListNode[Value]))的
|
||
# 继承字段展平需要从泛型基类的 stub 中读取字段(如 GSListNode.Next)。
|
||
# _generate_class_decl 内部会跳过泛型类的方法声明(T 参数 → opaque struct)。
|
||
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' | 'ctype_marker' | 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'
|
||
# CType 及其子类(CChar/CInt/CVoid/CPtr 等)是"类型强转"标记,
|
||
# 不应被编译为真实结构体,也不应被生成 stub 声明。
|
||
if issubclass(t_cls, t.CType):
|
||
return 'ctype_marker'
|
||
|
||
# 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'
|
||
# CType 及其子类是"类型强转"标记,不编译为结构体。
|
||
return 'ctype_marker'
|
||
|
||
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
|
||
elif kind == 'ctype_marker':
|
||
# CType 及其子类(CChar/CInt/CVoid/CPtr/CTypeDefault 等)是
|
||
# "类型强转"标记,不应被编译为真实结构体,也不生成 stub 声明。
|
||
return decls
|
||
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, <max_variant_payload> }
|
||
# 对齐 _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)
|
||
# PEP 695 泛型类跳过方法声明生成:方法参数类型 T 会被解析为 opaque struct,
|
||
# LLVM 报 "invalid type for function argument"。字段(AnnAssign)仍正常生成,
|
||
# 确保子类(如 Value(GSListNode[Value]))能从 stub 中读取继承字段(如 Next)。
|
||
IsGenericClass: bool = hasattr(node, 'type_params') and bool(node.type_params)
|
||
if not IsGenericClass:
|
||
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
|