524 lines
35 KiB
Python
524 lines
35 KiB
Python
from __future__ import annotations
|
||
from typing import Any
|
||
import ast
|
||
import os
|
||
import re
|
||
import llvmlite.ir as ir
|
||
from lib.core.Handles.HandlesBase import CTypeInfo
|
||
from lib.core.VLogger import get_logger as _vlog
|
||
from lib.constants.config import mode as _config_mode
|
||
|
||
|
||
class FuncGenMixin:
|
||
|
||
def setup_from_symbol_table(self, SymbolTable: Any) -> None:
|
||
self.SymbolTable: Any = SymbolTable
|
||
for name, info in SymbolTable.items():
|
||
if isinstance(info, CTypeInfo):
|
||
# CTypeInfo 对象:struct/union 由 StructGen 处理,此处只处理 dict 格式
|
||
continue
|
||
if not isinstance(info, dict):
|
||
if hasattr(info, 'get'):
|
||
info = dict(info) if hasattr(info, '__iter__') else {}
|
||
else:
|
||
continue
|
||
TypeKind: str = info.get('type', '')
|
||
if TypeKind == 'struct':
|
||
ClassName: str = name
|
||
self.class_methods[ClassName] = []
|
||
self.class_members[ClassName] = []
|
||
members: dict[str, Any] = info.get('members', {})
|
||
if members:
|
||
for MemberName, MemberInfo in members.items():
|
||
if isinstance(MemberInfo, dict):
|
||
MemberType: ir.Type = self._type_str_to_llvm(MemberInfo.get('type', 'int'), MemberInfo.get('IsPtr', False))
|
||
self.class_members[ClassName].append((MemberName, MemberType))
|
||
elif isinstance(MemberInfo, CTypeInfo):
|
||
MemberType: ir.Type = self._ctype_to_llvm(MemberInfo)
|
||
if MemberType is None:
|
||
MemberType = self._type_str_to_llvm(MemberInfo.ToString(), MemberInfo.IsPtr)
|
||
self.class_members[ClassName].append((MemberName, MemberType))
|
||
self._generate_structs()
|
||
self._create_Vtable_globals()
|
||
|
||
def register_method(self, ClassName: str, MethodName: str) -> None:
|
||
if ClassName not in self.class_methods:
|
||
self.class_methods[ClassName] = []
|
||
if ClassName not in self.class_members:
|
||
self.class_members[ClassName] = []
|
||
self._generate_structs()
|
||
self._create_Vtable_globals()
|
||
# 构造函数(__new__/__init__/__before_init__)不放入 vtable,
|
||
# 避免子类与基类 vtable 大小不一致导致虚分派索引偏移
|
||
_short_name: str = MethodName.split('.')[-1] if '.' in MethodName else MethodName
|
||
if _short_name in ('__new__', '__init__', '__before_init__'):
|
||
return
|
||
if MethodName not in self.class_methods[ClassName]:
|
||
self.class_methods[ClassName].append(MethodName)
|
||
|
||
def _apply_auto_addr(self, func_name: str, args: list[ir.Value]) -> list[ir.Value]:
|
||
"""对函数参数应用自动取地址逻辑(t.CNeedPtr / t.CAutoPtr)。
|
||
|
||
在调用前对标记为 auto-addr 的参数自动取地址:
|
||
- 'need' (t.CNeedPtr): 仅对非指针值取地址(alloca + store + get pointer)
|
||
- 'always' (t.CAutoPtr): 无条件取地址
|
||
"""
|
||
modes: list[str | None] | None = self._auto_addr_params.get(func_name)
|
||
if not modes:
|
||
return args
|
||
new_args: list[ir.Value] = list(args)
|
||
for i, mode in enumerate(modes):
|
||
if i >= len(new_args) or mode is None:
|
||
continue
|
||
arg: ir.Value = new_args[i]
|
||
is_ptr: bool = isinstance(arg.type, ir.PointerType)
|
||
if mode == 'always' or (mode == 'need' and not is_ptr):
|
||
# alloca + store + get pointer
|
||
slot: ir.Value = self.builder.alloca(arg.type)
|
||
self.builder.store(arg, slot)
|
||
new_args[i] = slot
|
||
return new_args
|
||
|
||
def _adjust_args(self, args: list[ir.Value], func: ir.Function) -> list[ir.Value]:
|
||
adjusted: list[ir.Value] = []
|
||
for i, (arg, param) in enumerate(zip(args, func.args)):
|
||
if arg.type != param.type:
|
||
if isinstance(arg.type, ir.PointerType) and isinstance(param.type, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)):
|
||
if arg.type.pointee == param.type or (hasattr(arg.type.pointee, 'name') and hasattr(param.type, 'name') and arg.type.pointee.name == param.type.name):
|
||
# C 语言中数组参数退化为指针,不做 load 而是退化为元素指针
|
||
if isinstance(param.type, ir.ArrayType):
|
||
zero: ir.Constant = ir.Constant(ir.IntType(32), 0)
|
||
elem_ptr: ir.Value = self.builder.gep(arg, [zero, zero], name=f"array_decay_arg_{i}")
|
||
adjusted.append(elem_ptr)
|
||
else:
|
||
adjusted.append(self._load(arg, name=f"Load_arg_{i}"))
|
||
continue
|
||
try:
|
||
if isinstance(arg.type, ir.IntType) and isinstance(param.type, ir.PointerType):
|
||
if isinstance(param.type.pointee, ir.IntType) and param.type.pointee.width == 8 and arg.type.width == 8:
|
||
# 当参数是 i8(字符值)而函数期望 i8*(字符串指针)时
|
||
# 为字符值分配内存,创建 null-terminated 字符串
|
||
tmp: ir.AllocaInstr = self._alloca(ir.ArrayType(ir.IntType(8), 2), name=f"char2str_arg_{i}")
|
||
zero: ir.Constant = ir.Constant(ir.IntType(32), 0)
|
||
char_ptr: ir.Value = self.builder.gep(tmp, [zero, zero], name=f"char2str_ptr_{i}")
|
||
self._store(arg, char_ptr)
|
||
null_ptr: ir.Value = self.builder.gep(tmp, [zero, ir.Constant(ir.IntType(32), 1)], name=f"char2str_null_{i}")
|
||
self._store(ir.Constant(ir.IntType(8), 0), null_ptr)
|
||
adjusted.append(char_ptr)
|
||
else:
|
||
if isinstance(arg.type, ir.PointerType) and isinstance(param.type, ir.IntType):
|
||
# 当参数是指针而函数期望整数时,使用 ptrtoint
|
||
adjusted.append(self.builder.ptrtoint(arg, param.type, name=f"ptrtoint_arg_{i}"))
|
||
else:
|
||
if arg.type.width < 32:
|
||
arg = self.builder.zext(arg, ir.IntType(32), name=f"zext_arg_{i}")
|
||
adjusted.append(self.builder.inttoptr(arg, param.type, name=f"inttoptr_arg_{i}"))
|
||
elif isinstance(arg.type, ir.PointerType) and isinstance(param.type, ir.IntType):
|
||
# 当参数是指针而函数期望整数时,使用 ptrtoint
|
||
adjusted.append(self.builder.ptrtoint(arg, param.type, name=f"ptrtoint_arg_{i}"))
|
||
elif isinstance(arg.type, ir.PointerType) and isinstance(param.type, ir.PointerType):
|
||
if isinstance(arg.type.pointee, ir.PointerType) and isinstance(param.type.pointee, ir.IntType):
|
||
Loaded: ir.Value = self._load(arg, name=f"Load_arg_{i}")
|
||
adjusted.append(Loaded)
|
||
elif isinstance(arg.type.pointee, ir.ArrayType) and isinstance(param.type.pointee, (ir.IntType, ir.FloatType, ir.DoubleType)):
|
||
zero: ir.Constant = ir.Constant(ir.IntType(32), 0)
|
||
elem_ptr: ir.Value = self.builder.gep(arg, [zero, zero], name=f"array_decay_arg_{i}")
|
||
if elem_ptr.type != param.type:
|
||
adjusted.append(self.builder.bitcast(elem_ptr, param.type, name=f"cast_arg_{i}"))
|
||
else:
|
||
adjusted.append(elem_ptr)
|
||
elif isinstance(arg.type.pointee, ir.ArrayType) and isinstance(param.type.pointee, ir.PointerType):
|
||
zero: ir.Constant = ir.Constant(ir.IntType(32), 0)
|
||
elem_ptr: ir.Value = self.builder.gep(arg, [zero, zero], name=f"array_decay_arg_{i}")
|
||
if elem_ptr.type != param.type:
|
||
adjusted.append(self.builder.bitcast(elem_ptr, param.type, name=f"cast_arg_{i}"))
|
||
else:
|
||
adjusted.append(elem_ptr)
|
||
else:
|
||
adjusted.append(self.builder.bitcast(arg, param.type, name=f"cast_arg_{i}"))
|
||
else:
|
||
if isinstance(arg.type, ir.PointerType) and isinstance(param.type, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)):
|
||
if arg.type.pointee == param.type or (hasattr(arg.type.pointee, 'name') and hasattr(param.type, 'name') and arg.type.pointee.name == param.type.name):
|
||
adjusted.append(self._load(arg, name=f"Load_arg_{i}"))
|
||
elif isinstance(arg.type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)):
|
||
tmp: ir.AllocaInstr = self._alloca(param.type, name=f"temp_arg_{i}")
|
||
Loaded: ir.Value = self.builder.bitcast(arg, ir.PointerType(param.type), name=f"cast_arg_{i}")
|
||
Loaded_val: ir.Value = self._load(Loaded, name=f"Load_arg_{i}")
|
||
self._store(Loaded_val, tmp)
|
||
adjusted.append(self._load(tmp, name=f"reLoad_arg_{i}"))
|
||
else:
|
||
tmp: ir.AllocaInstr = self._alloca(param.type, name=f"temp_arg_{i}")
|
||
cast_ptr: ir.Value = self.builder.bitcast(arg, ir.PointerType(param.type), name=f"cast_arg_{i}")
|
||
Loaded_val: ir.Value = self._load(cast_ptr, name=f"Load_arg_{i}")
|
||
self._store(Loaded_val, tmp)
|
||
adjusted.append(self._load(tmp, name=f"reLoad_arg_{i}"))
|
||
elif isinstance(arg.type, ir.PointerType) and arg.type.pointee == param.type:
|
||
adjusted.append(self._load(arg, name=f"Load_arg_{i}"))
|
||
elif isinstance(param.type, ir.PointerType) and param.type.pointee == arg.type:
|
||
tmp: ir.AllocaInstr = self._alloca(arg.type, name=f"temp_arg_{i}")
|
||
self._store(arg, tmp)
|
||
adjusted.append(tmp)
|
||
elif isinstance(arg.type, ir.IntType) and isinstance(param.type, ir.IntType):
|
||
if arg.type.width < param.type.width:
|
||
# 整数扩展:使用 sext(符号扩展)以正确处理负数
|
||
adjusted.append(self.builder.sext(arg, param.type, name=f"sext_arg_{i}"))
|
||
else:
|
||
adjusted.append(self.builder.trunc(arg, param.type, name=f"trunc_arg_{i}"))
|
||
elif isinstance(arg.type, (ir.FloatType, ir.DoubleType)) and isinstance(param.type, (ir.FloatType, ir.DoubleType)):
|
||
if isinstance(arg.type, ir.DoubleType) and isinstance(param.type, ir.FloatType):
|
||
adjusted.append(self.builder.fptrunc(arg, param.type, name=f"fptrunc_arg_{i}"))
|
||
elif isinstance(arg.type, ir.FloatType) and isinstance(param.type, ir.DoubleType):
|
||
adjusted.append(self.builder.fpext(arg, param.type, name=f"fpext_arg_{i}"))
|
||
else:
|
||
adjusted.append(arg)
|
||
elif isinstance(arg.type, (ir.FloatType, ir.DoubleType)) and isinstance(param.type, ir.IntType):
|
||
adjusted.append(self.builder.fptosi(arg, param.type, name=f"fptosi_arg_{i}"))
|
||
elif isinstance(arg.type, ir.IntType) and isinstance(param.type, (ir.FloatType, ir.DoubleType)):
|
||
adjusted.append(self.builder.sitofp(arg, param.type, name=f"sitofp_arg_{i}"))
|
||
elif isinstance(arg.type, ir.IntType) and isinstance(param.type, ir.PointerType):
|
||
# 整数转指针:inttoptr
|
||
if arg.type.width < 64:
|
||
i64_val: ir.Value = self.builder.sext(arg, ir.IntType(64), name=f"sext_arg_{i}")
|
||
adjusted.append(self.builder.inttoptr(i64_val, param.type, name=f"inttoptr_arg_{i}"))
|
||
else:
|
||
adjusted.append(self.builder.inttoptr(arg, param.type, name=f"inttoptr_arg_{i}"))
|
||
elif isinstance(arg.type, ir.PointerType) and isinstance(param.type, ir.IntType):
|
||
# 指针转整数:ptrtoint
|
||
i64_val: ir.Value = self.builder.ptrtoint(arg, ir.IntType(64), name=f"ptrtoint_arg_{i}")
|
||
if i64_val.type.width > param.type.width:
|
||
adjusted.append(self.builder.trunc(i64_val, param.type, name=f"trunc_arg_{i}"))
|
||
elif i64_val.type.width < param.type.width:
|
||
adjusted.append(self.builder.sext(i64_val, param.type, name=f"sext_arg_{i}"))
|
||
else:
|
||
adjusted.append(i64_val)
|
||
else:
|
||
adjusted.append(self.builder.bitcast(arg, param.type, name=f"cast_arg_{i}"))
|
||
except Exception: # 参数类型调整失败时使用简化回退逻辑
|
||
if isinstance(arg.type, ir.PointerType) and isinstance(param.type, (ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)):
|
||
if arg.type.pointee == param.type or (hasattr(arg.type.pointee, 'name') and hasattr(param.type, 'name') and arg.type.pointee.name == param.type.name):
|
||
adjusted.append(self._load(arg, name=f"Load_arg_{i}"))
|
||
else:
|
||
adjusted.append(self.builder.bitcast(arg, param.type, name=f"cast_arg_{i}"))
|
||
elif isinstance(arg.type, ir.PointerType) and arg.type.pointee == param.type:
|
||
adjusted.append(self._load(arg, name=f"Load_arg_{i}"))
|
||
elif isinstance(param.type, ir.PointerType) and param.type.pointee == arg.type:
|
||
tmp: ir.AllocaInstr = self._alloca(arg.type, name=f"temp_arg_{i}")
|
||
self._store(arg, tmp)
|
||
adjusted.append(tmp)
|
||
elif isinstance(arg.type, ir.PointerType) and isinstance(param.type, ir.PointerType):
|
||
# 任意指针到指针的类型转换(如 struct* -> i8*)
|
||
adjusted.append(self.builder.bitcast(arg, param.type, name=f"cast_arg_{i}"))
|
||
else:
|
||
adjusted.append(arg)
|
||
else:
|
||
adjusted.append(arg)
|
||
while len(adjusted) < len(func.args):
|
||
missing_idx: int = len(adjusted)
|
||
param: ir.Argument = func.args[missing_idx]
|
||
param_type: ir.Type = param.type if hasattr(param, 'type') else param
|
||
if isinstance(param_type, ir.PointerType):
|
||
adjusted.append(ir.Constant(param_type, None))
|
||
elif isinstance(param_type, ir.IntType):
|
||
adjusted.append(ir.Constant(param_type, 0))
|
||
elif isinstance(param_type, (ir.FloatType, ir.DoubleType)):
|
||
adjusted.append(ir.Constant(param_type, 0.0))
|
||
else:
|
||
adjusted.append(ir.Constant(ir.IntType(32), 0))
|
||
if func.type.pointee.var_arg:
|
||
for i in range(len(func.args), len(args)):
|
||
arg: ir.Value = args[i]
|
||
adjusted.append(arg)
|
||
return adjusted
|
||
|
||
def _get_member_offset(self, field_name: str, ClassName: str | None = None) -> int | None:
|
||
# 治本修复:同名类 sha1 冲突检测。
|
||
# 当 class_sha1_map[ClassName] 的 sha1 与 Gen.structs[ClassName] 的 sha1 不一致时,
|
||
# class_members[ClassName] 可能存储了错误模块的同名类成员。
|
||
# 此处检测冲突,加载正确模块的 class_members 并以 sha1 前缀全名缓存。
|
||
_ResolvedClass: str = ClassName if ClassName else ''
|
||
if ClassName:
|
||
_ClassSha1Map: dict[str, str] = getattr(self, 'class_sha1_map', {})
|
||
_ExpectedSha1: str | None = _ClassSha1Map.get(ClassName)
|
||
if _ExpectedSha1:
|
||
_ExistingStruct = self.structs.get(ClassName)
|
||
if isinstance(_ExistingStruct, ir.IdentifiedStructType):
|
||
_ExistingSha1: str = self._extract_struct_sha1(_ExistingStruct)
|
||
if _ExistingSha1 and _ExistingSha1 != _ExpectedSha1:
|
||
# sha1 冲突:使用全名查找 class_members
|
||
_FullKey: str = f"{_ExpectedSha1}.{ClassName}"
|
||
if _FullKey in self.class_members and self.class_members[_FullKey]:
|
||
_ResolvedClass = _FullKey
|
||
else:
|
||
# 从正确 stub 加载 class_members
|
||
if self._Trans and hasattr(self._Trans, 'ImportHandler'):
|
||
_StubName: str = f"{_ExpectedSha1}.stub.ll"
|
||
# 临时保存旧值,清除短名条目以绕过 guard,加载后恢复
|
||
_OldMembers = self.class_members.get(ClassName)
|
||
self.class_members[ClassName] = []
|
||
try:
|
||
self._Trans.ImportHandler._TryLoadClassMembersFromPyi(ClassName, _StubName, self)
|
||
except Exception:
|
||
pass
|
||
_NewMembers = self.class_members.get(ClassName)
|
||
if _NewMembers:
|
||
self.class_members[_FullKey] = _NewMembers
|
||
_ResolvedClass = _FullKey
|
||
# 恢复旧值(避免影响其他代码对短名的依赖)
|
||
if _OldMembers:
|
||
self.class_members[ClassName] = _OldMembers
|
||
else:
|
||
self.class_members.pop(ClassName, None)
|
||
has_vtable: bool = ClassName and (ClassName in self.class_vtable or ClassName in self._cross_module_vtable_classes)
|
||
base: int = 1 if has_vtable else 0
|
||
if ClassName and ClassName in self.structs:
|
||
struct_type: ir.IdentifiedStructType = self.structs[ClassName]
|
||
if isinstance(struct_type, ir.IdentifiedStructType) and struct_type.elements is not None and len(struct_type.elements) > 0:
|
||
if ClassName in self.class_members:
|
||
expected_fields: int = len(self.class_members[ClassName])
|
||
actual_elements: int = len(struct_type.elements)
|
||
if base > 0 and actual_elements == expected_fields:
|
||
base = 0
|
||
elif base == 0 and actual_elements == expected_fields + 1:
|
||
# 跨模块类可能未被检测为 CVTable(如 Name 继承 AST@t.CVTable),
|
||
# 但实际 struct 首元素是 vtable ptr → 需 base=1。
|
||
# 但若类有继承(class_parent 有值),多出的 1 可能是继承字段而非 vtable,
|
||
# 需检查 struct 首元素类型:i8*/i8** 为 vtable,其他为继承字段。
|
||
_first_elem = struct_type.elements[0] if actual_elements > 0 else None
|
||
_has_parent = self.class_parent.get(ClassName) is not None
|
||
_is_vtable_ptr = (isinstance(_first_elem, ir.PointerType) and
|
||
isinstance(_first_elem.pointee, ir.IntType) and _first_elem.pointee.width == 8)
|
||
if not _has_parent or _is_vtable_ptr:
|
||
base = 1
|
||
if _ResolvedClass and _ResolvedClass in self.class_members:
|
||
for i, (name, _) in enumerate(self.class_members[_ResolvedClass]):
|
||
if name == field_name:
|
||
return base + i
|
||
# 继承链回退:当字段不在当前类的 class_members 中时(常见于跨模块导入
|
||
# 父类未加载,字段被填充为 __inherited_N 占位名),沿 class_parent 链向上查找。
|
||
# 子类 struct 布局为 [vtable?] + [基类字段...] + [子类字段...],
|
||
# 基类字段在子类 struct 中的起始偏移即为 base(子类),
|
||
# 因此祖先 class_members 中的索引 i 对应子类 struct 中的偏移 base + i。
|
||
if ClassName and field_name:
|
||
_parent: str | None = self.class_parent.get(ClassName)
|
||
_visited: set = {ClassName} if ClassName else set()
|
||
while _parent and _parent not in _visited:
|
||
_visited.add(_parent)
|
||
# 跳过 __inherited_N 占位条目,只在真实字段名中查找
|
||
if _parent in self.class_members:
|
||
_members: list = self.class_members[_parent]
|
||
_has_placeholder: bool = any(n.startswith('__inherited_') for n, _ in _members)
|
||
if not _has_placeholder:
|
||
for i, (name, _) in enumerate(_members):
|
||
if name == field_name:
|
||
return base + i
|
||
# 父类 class_members 未加载或含占位条目,尝试从 stub 加载真实字段
|
||
if self._Trans and hasattr(self._Trans, 'ImportHandler'):
|
||
sha1_map: dict[str, str] = getattr(self, 'ModuleSha1Map', {})
|
||
for mod_name, mod_sha1 in sha1_map.items():
|
||
stub_name: str = f"{mod_sha1}.stub.ll"
|
||
self._Trans.ImportHandler._TryLoadClassMembersFromPyi(_parent, stub_name, self)
|
||
if _parent in self.class_members and not any(n.startswith('__inherited_') for n, _ in self.class_members[_parent]):
|
||
for i, (name, _) in enumerate(self.class_members[_parent]):
|
||
if name == field_name:
|
||
return base + i
|
||
break
|
||
_parent = self.class_parent.get(_parent)
|
||
short_name: str | None = ClassName.split('.')[-1] if ClassName and '.' in ClassName else None
|
||
if short_name and short_name in self.class_members:
|
||
for i, (name, _) in enumerate(self.class_members[short_name]):
|
||
if name == field_name:
|
||
return base + i
|
||
if ClassName and field_name:
|
||
if ClassName not in self.class_members or field_name not in [n for n, _ in self.class_members.get(ClassName, [])]:
|
||
if self._Trans and hasattr(self._Trans, 'ImportHandler'):
|
||
sha1_map: dict[str, str] = getattr(self, 'ModuleSha1Map', {})
|
||
for mod_name, mod_sha1 in sha1_map.items():
|
||
stub_name: str = f"{mod_sha1}.stub.ll"
|
||
self._Trans.ImportHandler._TryLoadClassMembersFromPyi(ClassName, stub_name, self)
|
||
if ClassName in self.class_members and len(self.class_members[ClassName]) > 0:
|
||
for i, (name, _) in enumerate(self.class_members[ClassName]):
|
||
if name == field_name:
|
||
return base + i
|
||
break
|
||
if short_name:
|
||
for mod_name, mod_sha1 in sha1_map.items():
|
||
stub_name: str = f"{mod_sha1}.stub.ll"
|
||
self._Trans.ImportHandler._TryLoadClassMembersFromPyi(short_name, stub_name, self)
|
||
if short_name in self.class_members and len(self.class_members[short_name]) > 0:
|
||
for i, (name, _) in enumerate(self.class_members[short_name]):
|
||
if name == field_name:
|
||
return base + i
|
||
break
|
||
if ClassName and field_name:
|
||
for cn_key in self.class_members:
|
||
if cn_key == ClassName or cn_key.endswith(f'.{ClassName}') or (short_name and (cn_key == short_name or cn_key.endswith(f'.{short_name}'))):
|
||
for i, (name, _) in enumerate(self.class_members[cn_key]):
|
||
if name == field_name:
|
||
return base + i
|
||
# REnum 变体搜索: ClassName 是 REnum(如 MetaType),成员注册在变体结构
|
||
# (MetaType_Var, MetaType_Concrete, MetaType_App)中。变体与 REnum 主结构
|
||
# 共享 max_variant_struct 布局,偏移一致(__tag 在偏移 0,字段从偏移 1 开始)
|
||
if ClassName and field_name:
|
||
variant_prefix: str = f"{ClassName}_"
|
||
for cn_key in self.class_members:
|
||
if cn_key.startswith(variant_prefix):
|
||
for i, (name, _) in enumerate(self.class_members[cn_key]):
|
||
if name == field_name:
|
||
return i
|
||
if short_name:
|
||
variant_prefix = f"{short_name}_"
|
||
for cn_key in self.class_members:
|
||
if cn_key.startswith(variant_prefix):
|
||
for i, (name, _) in enumerate(self.class_members[cn_key]):
|
||
if name == field_name:
|
||
return i
|
||
if ClassName and ClassName in self.structs:
|
||
struct_type: ir.IdentifiedStructType = self.structs[ClassName]
|
||
if isinstance(struct_type, ir.IdentifiedStructType) and struct_type.elements is not None and len(struct_type.elements) > 0:
|
||
for cn_key, members in self.class_members.items():
|
||
if cn_key == short_name or (short_name and cn_key.endswith(f'.{short_name}')):
|
||
for i, (name, _) in enumerate(members):
|
||
if name == field_name and i < len(struct_type.elements):
|
||
return i
|
||
break
|
||
else:
|
||
if short_name:
|
||
for cn_key, members in self.class_members.items():
|
||
sn: str = cn_key.split('.')[-1] if '.' in cn_key else cn_key
|
||
if sn == short_name:
|
||
for i, (name, _) in enumerate(members):
|
||
if name == field_name and i < len(struct_type.elements):
|
||
return i
|
||
break
|
||
if ClassName and field_name and ClassName not in self.class_members:
|
||
if self._Trans and hasattr(self._Trans, 'ImportHandler'):
|
||
for temp_dir_candidate in [getattr(self, '_temp_dir', None)]:
|
||
if temp_dir_candidate and os.path.isdir(temp_dir_candidate):
|
||
for pyi_file in os.listdir(temp_dir_candidate):
|
||
if pyi_file.endswith('.pyi'):
|
||
stub_name: str = pyi_file.replace('.pyi', '.stub.ll')
|
||
short_cn: str = short_name if short_name else ClassName
|
||
self._Trans.ImportHandler._TryLoadClassMembersFromPyi(short_cn, stub_name, self)
|
||
if short_cn in self.class_members and len(self.class_members[short_cn]) > 0:
|
||
for i, (name, _) in enumerate(self.class_members[short_cn]):
|
||
if name == field_name:
|
||
return base + i
|
||
break
|
||
if ClassName and ClassName in self.structs and (ClassName not in self.class_members or len(self.class_members.get(ClassName, [])) == 0):
|
||
struct_type: ir.IdentifiedStructType = self.structs[ClassName]
|
||
if isinstance(struct_type, ir.IdentifiedStructType) and struct_type.elements is not None:
|
||
if self._Trans and hasattr(self._Trans, 'ClassHandler'):
|
||
class_handler: Any = self._Trans.ClassHandler
|
||
if self._current_tree:
|
||
for node in ast.iter_child_nodes(self._current_tree):
|
||
if isinstance(node, ast.ClassDef) and (node.name == ClassName or node.name == short_name):
|
||
if node.name not in self.class_members:
|
||
self.class_members[node.name] = []
|
||
if len(self.class_members[node.name]) == 0:
|
||
class_handler._PreRegisterClassMembers(node, self)
|
||
if ClassName in self.class_members:
|
||
for i, (name, _) in enumerate(self.class_members[ClassName]):
|
||
if name == field_name:
|
||
return base + i
|
||
if short_name and short_name in self.class_members:
|
||
for i, (name, _) in enumerate(self.class_members[short_name]):
|
||
if name == field_name:
|
||
return base + i
|
||
break
|
||
# 最后回退:当继承链回退失败但 struct body 已加载时,
|
||
# 从 struct elements 数和 class_members 数推断继承字段偏移。
|
||
# 继承字段在 struct 开头(base 之后),此 fallback 仅适用于单继承字段的情况
|
||
# (如 GSListNode[T] 的 Next 字段)。子类 struct 布局:
|
||
# [vtable?] + [基类字段...] + [子类字段...]
|
||
# 当 class_members 只含子类自身字段、基类 class_members 未加载时,
|
||
# 继承字段数 = struct_elements - class_members_count,
|
||
# 若该值 == 1,则所查字段(如 Next)的偏移 = base + 0。
|
||
if ClassName and field_name:
|
||
_st_final = self.structs.get(ClassName)
|
||
if isinstance(_st_final, ir.IdentifiedStructType) and _st_final.elements is not None:
|
||
_num_elements_final: int = len(_st_final.elements)
|
||
_members_final: list = self.class_members.get(ClassName, [])
|
||
_num_members_final: int = len(_members_final)
|
||
if _num_elements_final > _num_members_final:
|
||
_num_inherited_final: int = _num_elements_final - _num_members_final
|
||
if _num_inherited_final == 1:
|
||
# 单继承字段,偏移 = base(vtable 之后第一个位置)
|
||
return base + 0
|
||
return None
|
||
|
||
def _resolve_class_for_var(self, var_name: str) -> str | None:
|
||
for ClassName in self.class_methods:
|
||
if var_name == ClassName or var_name.startswith(f"__with_{ClassName}_"):
|
||
return ClassName
|
||
if var_name in self.variables:
|
||
var: ir.Value = self.variables[var_name]
|
||
if isinstance(var.type, ir.PointerType):
|
||
for ClassName, struct_type in self.structs.items():
|
||
if var.type.pointee == struct_type:
|
||
return ClassName
|
||
if isinstance(var.type.pointee, ir.PointerType) and var.type.pointee.pointee == struct_type:
|
||
return ClassName
|
||
return None
|
||
|
||
def _infer_return_type(self, FuncName: str) -> ir.Type:
|
||
if FuncName in self.functions:
|
||
return self.functions[FuncName].type.pointee.return_type
|
||
# 也尝试混淆后的名称
|
||
mangled: str = self._mangle_func_name(FuncName)
|
||
if mangled != FuncName and mangled in self.functions:
|
||
return self.functions[mangled].type.pointee.return_type
|
||
if FuncName in self.known_return_types:
|
||
return self.known_return_types[FuncName]
|
||
if mangled != FuncName and mangled in self.known_return_types:
|
||
return self.known_return_types[mangled]
|
||
for ClassName, methods in self.class_methods.items():
|
||
for method in methods:
|
||
if FuncName == method:
|
||
return ir.VoidType()
|
||
if self._import_handler_ref:
|
||
try:
|
||
# 先用原始名称查找,再用混淆名称查找
|
||
stub_func_type: Any = self._import_handler_ref._LookupStubFuncType(FuncName, self)
|
||
if stub_func_type is None and mangled != FuncName:
|
||
stub_func_type = self._import_handler_ref._LookupStubFuncType(mangled, self)
|
||
if stub_func_type and hasattr(stub_func_type, 'return_type'):
|
||
return stub_func_type.return_type
|
||
except Exception as _e:
|
||
if _config_mode == "strict":
|
||
raise
|
||
_vlog().warning(f"查找桩函数返回类型失败: {_e}", "Exception")
|
||
return ir.IntType(32)
|
||
|
||
def _get_or_declare_func(self, name: str, func_type: ir.FunctionType) -> ir.Function:
|
||
for fname, fobj in self.functions.items():
|
||
if fname == name:
|
||
return fobj
|
||
if re.match(r'^[0-9a-f]{16,}\.' + re.escape(name) + r'$', fname):
|
||
return fobj
|
||
func_decl: ir.Function = ir.Function(self.module, func_type, name=name)
|
||
self.functions[name] = func_decl
|
||
return func_decl
|
||
|
||
def get_or_declare_c_func(self, name: str, fallback_type: ir.FunctionType | None = None) -> ir.Function | None:
|
||
if name in self.functions:
|
||
return self.functions[name]
|
||
for fname, fobj in self.functions.items():
|
||
if re.match(r'^[0-9a-f]{16,}\.' + re.escape(name) + r'$', fname):
|
||
return fobj
|
||
stub_func_type: ir.FunctionType | None = None
|
||
if self._import_handler_ref:
|
||
try:
|
||
stub_func_type = self._import_handler_ref._LookupStubFuncType(name, self)
|
||
except Exception as _e:
|
||
if _config_mode == "strict":
|
||
raise
|
||
_vlog().warning(f"查找桩函数类型失败: {_e}", "Exception")
|
||
if stub_func_type:
|
||
func_decl: ir.Function = ir.Function(self.module, stub_func_type, name=name)
|
||
self.functions[name] = func_decl
|
||
return func_decl
|
||
if fallback_type is not None:
|
||
func_decl: ir.Function = ir.Function(self.module, fallback_type, name=name)
|
||
self.functions[name] = func_decl
|
||
return func_decl
|
||
return None |