Files
TransPyC/lib/core/LLVMCG/FuncGen.py

437 lines
28 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
from typing import 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()
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)
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:
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
if ClassName and ClassName in self.class_members:
for i, (name, _) in enumerate(self.class_members[ClassName]):
if name == field_name:
return base + i
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 hasattr(self, '_Trans') and 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
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 hasattr(self, '_Trans') and 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 hasattr(self, '_Trans') and self._Trans and hasattr(self._Trans, 'ClassHandler'):
class_handler: Any = self._Trans.ClassHandler
if hasattr(self, '_current_tree') and 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
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 _GetOrCreateFunc(self, FuncName: str, ArgTypes: list[ir.Type]) -> ir.Function:
mangled_name: str = self._mangle_func_name(FuncName)
if mangled_name in self.functions:
return self.functions[mangled_name]
FuncType: ir.FunctionType = self._InferFuncTypeForCall(FuncName, ArgTypes)
try:
func: ir.Function = ir.Function(self.module, FuncType, name=mangled_name)
except Exception: # 函数名冲突时添加前缀重试
func = ir.Function(self.module, FuncType, name=f"_{mangled_name}")
self.functions[mangled_name] = func
return func
def _resolve_method_class(self, MethodName: str) -> str | None:
for ClassName, methods in self.class_methods.items():
if MethodName in methods:
return ClassName
full_name: str = f"{ClassName}.{MethodName}"
if full_name in methods:
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 hasattr(self, '_import_handler_ref') and 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 _InferFuncTypeForCall(self, FuncName: str, ArgTypes: list[ir.Type]) -> ir.FunctionType:
return_type: ir.Type = self._infer_return_type(FuncName)
if not ArgTypes:
for ClassName, methods in self.class_methods.items():
if FuncName in methods:
if ClassName in self.structs:
ArgTypes = [ir.PointerType(self.structs[ClassName])]
else:
ArgTypes = [ir.PointerType(ir.LiteralStructType([ir.PointerType(ir.IntType(8)), ir.IntType(32)]))]
break
else:
ArgTypes = [ir.IntType(32)]
return ir.FunctionType(return_type, ArgTypes)
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 hasattr(self, '_import_handler_ref') and 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