Files
TransPyC/lib/core/Handles/HandlesImports.py

2016 lines
112 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 TYPE_CHECKING
if TYPE_CHECKING:
from lib.core.translator import Translator
from lib.core.Translator.LlvmGenerator import LlvmGeneratorMixin
import ast
import inspect
import os
import re
import hashlib
import llvmlite.ir as ir
from lib.core.Handles.HandlesBase import BaseHandle, CTypeInfo, FuncMeta
from lib.core.VLogger import get_logger as _vlog
from lib.constants.config import mode as _config_mode
from lib.includes import t
from lib.includes.t import CTypeRegistry
from lib.core.SymbolUtils import IsTModule, ParseListAnnotation, ExtractTypeNameFromBinOp, AnnotationContainsName
class ImportHandle(BaseHandle):
_stub_cache = {}
_stub_cache_dir = None
_struct_Load_cache = {}
_pyi_cache = {}
_pyi_cache_dir = None
_project_root_cache = None
def _reset_file_cache(self) -> None:
self._struct_Load_cache.clear()
def _find_project_root(self) -> str:
if self._project_root_cache is not None:
return self._project_root_cache
# 优先使用 translator 上显式设置的 _temp_dir由 Phase2Translator._translate_file 传入),
# 这样可以正确处理 include 文件——其源文件目录不在项目根目录下,
# 从源文件目录向上搜索 temp/output 会失败并回退到错误的路径。
temp_dir_attr = getattr(self.Trans, '_temp_dir', None)
if temp_dir_attr and os.path.isdir(temp_dir_attr):
project_root = os.path.dirname(os.path.abspath(temp_dir_attr))
if os.path.isdir(os.path.join(project_root, 'temp')) and os.path.isdir(os.path.join(project_root, 'output')):
self._project_root_cache = project_root
return project_root
current_file = getattr(self.Trans, 'CurrentFile', '') or ''
if current_file:
current_dir = os.path.dirname(os.path.abspath(current_file))
search_dir = current_dir
for _ in range(10):
if os.path.isdir(os.path.join(search_dir, 'temp')) and os.path.isdir(os.path.join(search_dir, 'output')):
self._project_root_cache = search_dir
return search_dir
parent = os.path.dirname(search_dir)
if parent == search_dir:
break
search_dir = parent
# 回退到当前工作目录(编译运行时的项目根目录)
cwd = os.getcwd()
if os.path.isdir(os.path.join(cwd, 'temp')) and os.path.isdir(os.path.join(cwd, 'output')):
self._project_root_cache = cwd
return cwd
if current_file:
self._project_root_cache = current_dir
return current_dir
self._project_root_cache = cwd
return cwd
def _EmitImportDeclarationsLlvm(self, Node: ast.Import, Gen: LlvmGeneratorMixin) -> None:
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
for alias in Node.names:
name = alias.name
if name in ('c', 't'):
if alias.asname:
self.Trans.SymbolTable.import_aliases[alias.asname] = name
AliasInfo = CTypeInfo()
AliasInfo.IsModuleAlias = True
AliasInfo.ResolvedModule = name
self.Trans.SymbolTable.insert(alias.asname, AliasInfo)
continue
self.Trans._ImportedModules.add(name)
if alias.asname:
self.Trans.SymbolTable.import_aliases[alias.asname] = name
# 同时在符号表中注册模块别名,以便 CTypeInfo.FromNode 能解析
AliasInfo = CTypeInfo()
AliasInfo.IsModuleAlias = True
AliasInfo.ResolvedModule = name
self.Trans.SymbolTable.insert(alias.asname, AliasInfo)
current_module = self._get_current_module_name()
self._EmitModuleDeclarationsLlvm(name, Gen, register_module_name=current_module)
def _RegisterFromImportAliases(self, Node: ast.ImportFrom, Gen: LlvmGeneratorMixin, module: str) -> None:
"""Register 'from module import y as z' aliases after module declarations are Loaded"""
if not Node.names:
return
for alias in Node.names:
name = alias.name
asname = alias.asname
# For 'from X import Y' (no asname), register CDefine constants
# into _define_constants so they can be found by _HandleNameLlvm
if not asname or asname == name:
# Check if this name is a CDefine constant in SymbolTable
sym_info = self.Trans.SymbolTable.lookup(name)
if sym_info and sym_info.IsDefine and sym_info.DefineValue is not None:
define_constants = vars(Gen).setdefault('_define_constants', {})
if name not in define_constants:
define_constants[name] = sym_info.DefineValue
# Also check with module prefix
for prefix_key in [f"{module}.{name}", name]:
sym_info2 = self.Trans.SymbolTable.lookup(prefix_key)
if sym_info2 and sym_info2.IsDefine and sym_info2.DefineValue is not None:
define_constants = vars(Gen).setdefault('_define_constants', {})
if name not in define_constants:
define_constants[name] = sym_info2.DefineValue
break
continue
if not hasattr(self.Trans, '_t_c_imported_names'):
self.Trans._t_c_imported_names = {}
self.Trans._t_c_imported_names[asname] = (module, name)
if name in Gen.functions:
Gen.functions[asname] = Gen.functions[name]
if name in Gen.structs:
Gen.structs[asname] = Gen.structs[name]
info = self.Trans.SymbolTable.lookup(name)
if info:
self.Trans.SymbolTable.insert(asname, info)
if name in Gen.variables:
Gen.variables[asname] = Gen.variables[name]
if name in Gen.global_vars:
Gen.global_vars[asname] = Gen.global_vars[name]
for func_name in list(Gen.functions.keys()):
if func_name.startswith(f'{name}.'):
method_suffix = func_name[len(name):]
Gen.functions[f'{asname}{method_suffix}'] = Gen.functions[func_name]
for key in list(self.Trans.SymbolTable.keys()):
if key.startswith(f'{name}.'):
suffix = key[len(name):]
self.Trans.SymbolTable.insert(f'{asname}{suffix}', self.Trans.SymbolTable[key])
for meta_dict in (Gen.class_members, Gen.class_member_defaults,
Gen.class_member_signeds, Gen.class_member_bitfields,
Gen.class_member_byteorders, Gen.class_member_bitoffsets,
Gen.class_methods):
if name in meta_dict:
meta_dict[asname] = meta_dict[name]
if name in Gen.class_vtable:
Gen.class_vtable.add(asname)
if name in Gen.class_packed:
Gen.class_packed.add(asname)
struct_sha1_map = getattr(Gen, '_struct_sha1_map', {})
if name in struct_sha1_map:
struct_sha1_map[asname] = struct_sha1_map[name]
def _EmitImportFromDeclarationsLlvm(self, Node: ast.ImportFrom, Gen: LlvmGeneratorMixin) -> None:
module = Node.module if Node.module else ''
if module in ('c', 't'):
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(module)
if not hasattr(self.Trans, '_t_c_imported_names'):
self.Trans._t_c_imported_names = {}
# 注册 t/c 类型别名到 _t_type_symbolsCType 子类)
t_type_syms = self.Trans.SymbolTable._t_type_symbols
for alias in Node.names:
name = alias.name
asname = alias.asname or name
self.Trans._t_c_imported_names[asname] = (module, name)
# 仅注册 CType 子类CEnum/CUnion/CStruct/REnum 等)
if module == 't':
cls = getattr(t, name, None)
if isinstance(cls, type) and issubclass(cls, t.CType):
t_type_syms[asname] = cls
return
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
current_module = self._get_current_module_name()
if Node.level and Node.level > 0:
current_file = getattr(self.Trans, 'CurrentFile', '') or ''
current_dir = os.path.dirname(os.path.abspath(current_file))
parent_dir = current_dir
for _ in range(Node.level - 1):
parent_dir = os.path.dirname(parent_dir)
if module:
SearchExtensions = ['.pyi', '.py']
sub_path_base = os.path.join(parent_dir, module)
for ext in SearchExtensions:
candidate = sub_path_base + ext
if os.path.isfile(candidate):
self.Trans._ImportedModules.add(module)
self._LoadModuleDeclarationsFromFile(candidate, Gen, module_name=module, register_module_name=current_module)
self._RegisterFromImportAliases(Node, Gen, module)
return
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
# 尝试完整模块名和短模块名
full_mod_name = f"{current_module}.{module}" if current_module else module
target_sha1 = ModuleSha1Map.get(full_mod_name) or ModuleSha1Map.get(module)
if target_sha1:
temp_dir = getattr(Gen, '_temp_dir', None)
if temp_dir:
sha1_pyi = os.path.join(temp_dir, f"{target_sha1}.pyi")
if os.path.isfile(sha1_pyi):
self.Trans._ImportedModules.add(module)
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=module, register_module_name=current_module)
self._RegisterFromImportAliases(Node, Gen, module)
return
else:
mod_dir = parent_dir
SearchExtensions = ['.pyi', '.py']
pkg_name = ''
if not module:
# When 'from . import name' is used, Load the package's __init__
# to make package-level names (functions, classes, constants) available
pkg_name = self._LoadPackageInitForRelativeImport(mod_dir, Gen, current_module)
# Also try to Load each alias as a submodule
for alias in Node.names:
found_sub = False
for ext in SearchExtensions:
sub_path = os.path.join(mod_dir, alias.name + ext)
if os.path.isfile(sub_path):
self.Trans._ImportedModules.add(alias.name)
self._LoadModuleDeclarationsFromFile(sub_path, Gen, module_name=alias.name, register_module_name=current_module)
if alias.asname:
self.Trans.SymbolTable.import_aliases[alias.asname] = alias.name
found_sub = True
break
if not found_sub:
# SHA1 map fallback for submodule
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
sub_ModulePath = f"{pkg_name}.{alias.name}" if pkg_name else alias.name
target_sha1 = ModuleSha1Map.get(sub_ModulePath) or ModuleSha1Map.get(alias.name)
if target_sha1:
temp_dir = getattr(Gen, '_temp_dir', None)
if temp_dir:
sha1_pyi = os.path.join(temp_dir, f"{target_sha1}.pyi")
if os.path.isfile(sha1_pyi):
self.Trans._ImportedModules.add(sub_ModulePath)
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=alias.name, register_module_name=current_module)
if alias.asname:
self.Trans.SymbolTable.import_aliases[alias.asname] = alias.name
self._RegisterFromImportAliases(Node, Gen, module or pkg_name)
return
self.Trans._ImportedModules.add(module)
self._EmitModuleDeclarationsLlvm(module, Gen, register_module_name=current_module)
self._RegisterFromImportAliases(Node, Gen, module)
def _LoadModuleDeclarationsFromFile(self, pyi_path: str, Gen: LlvmGeneratorMixin, module_name: str | None = None, register_module_name: str | None = None, actual_module_name: str | None = None, reexport_package: str | None = None) -> None:
if not pyi_path or not os.path.isfile(pyi_path):
return
try:
with open(pyi_path, 'r', encoding='utf-8') as f:
ModuleCode = f.read()
ModuleTree = ast.parse(ModuleCode)
for node in ModuleTree.body:
if isinstance(node, ast.FunctionDef):
if node.name in ('va_start', 'va_arg', 'va_end'):
continue
file_module_name = os.path.splitext(os.path.basename(pyi_path))[0]
effective_register = register_module_name or module_name
reexport_names = []
if module_name and module_name != file_module_name and module_name != effective_register:
reexport_names.append(module_name)
# When Loading a submodule for 'from .xxx import yyy' inside a package's __init__.py,
# re-export functions under the package name so 'pkg.func()' resolves correctly
if reexport_package and reexport_package != file_module_name and reexport_package not in reexport_names:
reexport_names.append(reexport_package)
self._EmitExternalFuncDeclLlvm(node, Gen, source_module_name=file_module_name, register_module_name=effective_register, reexport_module_names=reexport_names if reexport_names else None)
elif isinstance(node, ast.AnnAssign):
self._EmitExternalGlobalDeclLlvm(node, Gen, module_name=module_name, ModulePath=pyi_path)
elif isinstance(node, ast.Assign):
self._EmitExternalGlobalDeclLlvm(node, Gen, module_name=module_name, ModulePath=pyi_path)
elif isinstance(node, ast.ClassDef):
if hasattr(node, 'type_params') and node.type_params:
if not hasattr(self.Trans.ClassHandler, '_generic_class_templates'):
self.Trans.ClassHandler._generic_class_templates = {}
type_params = [tp.name for tp in node.type_params]
self.Trans.ClassHandler._generic_class_templates[node.name] = {
'node': node,
'type_params': type_params,
}
self._EmitExternalClassDeclLlvm(node, Gen, module_name=module_name, actual_module_name=actual_module_name)
elif isinstance(node, ast.ImportFrom):
if node.level > 0:
current_dir = os.path.dirname(os.path.abspath(pyi_path))
parent_dir = current_dir
for _ in range(node.level - 1):
parent_dir = os.path.dirname(parent_dir)
if node.module:
SearchExtensions = ['.pyi', '.py']
sub_path_base = os.path.join(parent_dir, node.module)
found_sub = False
for ext in SearchExtensions:
candidate = sub_path_base + ext
if os.path.isfile(candidate):
actual_module = f"{register_module_name}.{node.module}" if register_module_name else node.module
self._RegisterSubModuleSha1(candidate, actual_module, node.module, Gen)
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(actual_module)
for alias in node.names:
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(candidate, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=reexport_package or module_name)
else:
self._LoadModuleDeclarationsFromFile(candidate, Gen, module_name=alias.name, register_module_name=register_module_name, actual_module_name=actual_module, reexport_package=reexport_package or module_name)
found_sub = True
break
# SHA1 map 回退
if not found_sub:
actual_module = f"{register_module_name}.{node.module}" if register_module_name else node.module
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
target_sha1 = ModuleSha1Map.get(actual_module) or ModuleSha1Map.get(node.module)
if target_sha1:
temp_dir = getattr(Gen, '_temp_dir', None)
if temp_dir:
sha1_pyi = os.path.join(temp_dir, f"{target_sha1}.pyi")
if os.path.isfile(sha1_pyi):
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(actual_module)
for alias in node.names:
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=reexport_package or module_name)
else:
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=alias.name, register_module_name=register_module_name, actual_module_name=actual_module, reexport_package=reexport_package or module_name)
else:
mod_dir = parent_dir
SearchExtensions = ['.pyi', '.py']
# Load __init__.py to make package-level names available
self._LoadPackageInitForRelativeImport(mod_dir, Gen, register_module_name)
for alias in node.names:
found_alias = False
for ext in SearchExtensions:
sub_path = os.path.join(mod_dir, alias.name + ext)
if os.path.isfile(sub_path):
sub_ModulePath = f"{register_module_name}.{alias.name}" if register_module_name else alias.name
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(sub_ModulePath)
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(sub_path, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=reexport_package or module_name)
else:
self._LoadModuleDeclarationsFromFile(sub_path, Gen, module_name=alias.name, register_module_name=register_module_name, reexport_package=reexport_package or module_name)
found_alias = True
break
# SHA1 map 回退
if not found_alias:
sub_ModulePath = f"{register_module_name}.{alias.name}" if register_module_name else alias.name
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
target_sha1 = ModuleSha1Map.get(sub_ModulePath) or ModuleSha1Map.get(alias.name)
if target_sha1:
temp_dir = getattr(Gen, '_temp_dir', None)
if temp_dir:
sha1_pyi = os.path.join(temp_dir, f"{target_sha1}.pyi")
if os.path.isfile(sha1_pyi):
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(sub_ModulePath)
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=reexport_package or module_name)
else:
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=alias.name, register_module_name=register_module_name, reexport_package=reexport_package or module_name)
except Exception as _e:
if _config_mode == "strict":
self.Trans.LogWarning(f"异常被忽略: {_e}")
def _get_current_module_name(self) -> str | None:
current_file = getattr(self.Trans, 'CurrentFile', '') or ''
if current_file:
return os.path.splitext(os.path.basename(current_file))[0]
return None
def _RegisterSubModuleSha1(self, candidate_path: str, actual_module: str, short_module: str, Gen: LlvmGeneratorMixin) -> None:
"""Register a submodule's SHA1 in ModuleSha1Map so cross-module class references work"""
try:
with open(candidate_path, 'r', encoding='utf-8') as f:
sub_content = f.read()
sub_sha1 = hashlib.sha1(sub_content.encode('utf-8')).hexdigest()[:16]
if hasattr(Gen, 'ModuleSha1Map'):
Gen.ModuleSha1Map[actual_module] = sub_sha1
if short_module and short_module not in Gen.ModuleSha1Map:
Gen.ModuleSha1Map[short_module] = sub_sha1
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
def _LoadPackageInitForRelativeImport(self, mod_dir: str, Gen: LlvmGeneratorMixin, register_module_name: str | None) -> str:
"""Load __init__.py from the package directory for 'from . import name' resolution.
This makes package-level names (functions, classes, constants) available.
Returns the package name."""
pkg_name = os.path.basename(mod_dir)
init_Loaded = False
# 跳过当前正在编译的文件,避免重复加载导致 class_members 重复
current_sha1 = getattr(Gen, 'module_sha1', None)
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
pkg_sha1 = ModuleSha1Map.get(pkg_name)
if pkg_sha1 and pkg_sha1 == current_sha1:
return pkg_name
# Try SHA1 map first (correct mangled names from stub)
if pkg_sha1:
temp_dir = getattr(Gen, '_temp_dir', None)
if temp_dir:
sha1_pyi = os.path.join(temp_dir, f"{pkg_sha1}.pyi")
if os.path.isfile(sha1_pyi):
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=pkg_name, register_module_name=register_module_name)
init_Loaded = True
# Fallback: try Loading __init__.py directly from the package directory
if not init_Loaded:
for ext in ['.pyi', '.py']:
init_path = os.path.join(mod_dir, '__init__' + ext)
if os.path.isfile(init_path):
self._LoadModuleDeclarationsFromFile(init_path, Gen, module_name=pkg_name, register_module_name=register_module_name)
init_Loaded = True
break
return pkg_name
def _EmitModuleDeclarationsLlvm(self, module_name: str, Gen: LlvmGeneratorMixin, register_module_name: str | None = None) -> None:
if module_name in ('pyzlib',):
source_sig_files = getattr(self.Trans, '_source_module_sig_files', None)
ModulePath = module_name.replace('.', os.sep) + '.pyi'
PackageInitPath = module_name.replace('.', os.sep) + os.sep + '__init__.pyi'
ProjectRoot = os.path.dirname(os.path.abspath(getattr(self.Trans, 'CurrentFile', '') or '')) or os.getcwd()
HandlesFile = os.path.abspath(inspect.getfile(self.__class__))
TransPyCRoot = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(HandlesFile))))
FullModulePath = None
SearchDirs = [os.path.join(TransPyCRoot, 'includes'), os.path.join(TransPyCRoot, 'lib', 'includes')] + list(self.Trans.LibraryPaths)
for LibPath in SearchDirs:
SearchPath = os.path.join(ProjectRoot, LibPath) if not os.path.isabs(LibPath) else LibPath
CandidatePath = os.path.join(SearchPath, ModulePath)
if os.path.isfile(CandidatePath):
FullModulePath = CandidatePath
break
CandidatePath = os.path.join(SearchPath, PackageInitPath)
if os.path.isfile(CandidatePath):
FullModulePath = CandidatePath
break
ModuleParts = ModulePath.split(os.sep)
if len(ModuleParts) > 1:
LibPathBasename = os.path.basename(LibPath.rstrip('/\\'))
if LibPathBasename == ModuleParts[0]:
RemainingPath = os.sep.join(ModuleParts[1:])
CandidatePath = os.path.join(SearchPath, RemainingPath)
if os.path.isfile(CandidatePath):
FullModulePath = CandidatePath
break
if not FullModulePath:
ModulePathPy = module_name.replace('.', os.sep) + '.py'
PackageInitPathPy = module_name.replace('.', os.sep) + os.sep + '__init__.py'
SearchDirsForPy = [os.path.join(TransPyCRoot, 'includes'), os.path.join(TransPyCRoot, 'lib', 'includes')] + list(self.Trans.LibraryPaths)
for LibPath in SearchDirsForPy:
SearchPath = os.path.join(ProjectRoot, LibPath) if not os.path.isabs(LibPath) else LibPath
CandidatePath = os.path.join(SearchPath, ModulePathPy)
if os.path.isfile(CandidatePath):
FullModulePath = CandidatePath
break
CandidatePath = os.path.join(SearchPath, PackageInitPathPy)
if os.path.isfile(CandidatePath):
FullModulePath = CandidatePath
break
if not FullModulePath or not os.path.isfile(FullModulePath):
source_sig_files = getattr(self.Trans, '_source_module_sig_files', None)
if source_sig_files:
sig_path = source_sig_files.get(module_name)
if sig_path and os.path.isfile(sig_path):
FullModulePath = sig_path
else:
for key, val in source_sig_files.items():
if key.endswith('.' + module_name) or key == module_name:
FullModulePath = val
break
if not FullModulePath or not os.path.isfile(FullModulePath):
if module_name in ('pyzlib', 'zdeflate', 'zinflate', 'zchecksum', 'zdef', 'zhuff'):
pass
if not FullModulePath:
CurrentFile = getattr(self.Trans, 'CurrentFile', None) or ''
if CurrentFile:
CurrentFileDir = os.path.dirname(os.path.abspath(CurrentFile))
SearchFromDirs = [CurrentFileDir, os.getcwd()]
else:
SearchFromDirs = [os.getcwd()]
SearchDirsForSubmodule = SearchFromDirs + [os.path.join(TransPyCRoot, 'includes'), os.path.join(TransPyCRoot, 'lib', 'includes')] + list(self.Trans.LibraryPaths)
for SearchPath in SearchDirsForSubmodule:
if os.path.isdir(SearchPath):
for root, dirs, files in os.walk(SearchPath):
if '__pycache__' in root:
continue
if module_name + '.py' in files:
FullModulePath = os.path.join(root, module_name + '.py')
break
if module_name + '.pyi' in files:
FullModulePath = os.path.join(root, module_name + '.pyi')
break
if FullModulePath:
break
self._LoadDeclarationsFromStubLlvm(module_name, Gen)
if module_name == 'pyzlib':
pass
try:
with open(FullModulePath, 'r', encoding='utf-8') as f:
ModuleCode = f.read()
ModuleTree = ast.parse(ModuleCode)
if module_name == 'pyzlib':
func_names = [n.name for n in ModuleTree.body if isinstance(n, ast.FunctionDef)]
has_functions_or_classes = False
for node in ModuleTree.body:
if module_name == 'pyzlib':
node_type = type(node).__name__
name = getattr(node, 'name', getattr(node, 'attr', '?'))
try:
if isinstance(node, ast.FunctionDef):
if node.name in ('va_start', 'va_arg', 'va_end'):
continue
effective_register = register_module_name or module_name
self._EmitExternalFuncDeclLlvm(node, Gen, source_module_name=module_name, register_module_name=effective_register)
has_functions_or_classes = True
elif isinstance(node, ast.AnnAssign):
self._EmitExternalGlobalDeclLlvm(node, Gen, module_name=module_name, ModulePath=FullModulePath)
has_functions_or_classes = True
elif isinstance(node, ast.Assign):
self._EmitExternalGlobalDeclLlvm(node, Gen, module_name=module_name, ModulePath=FullModulePath)
has_functions_or_classes = True
elif isinstance(node, ast.ClassDef):
if hasattr(node, 'type_params') and node.type_params:
if not hasattr(self.Trans.ClassHandler, '_generic_class_templates'):
self.Trans.ClassHandler._generic_class_templates = {}
type_params = [tp.name for tp in node.type_params]
self.Trans.ClassHandler._generic_class_templates[node.name] = {
'node': node,
'type_params': type_params,
}
self._EmitExternalClassDeclLlvm(node, Gen, module_name=module_name)
has_functions_or_classes = True
elif isinstance(node, ast.ImportFrom):
if node.level > 0:
current_dir = os.path.dirname(os.path.abspath(FullModulePath))
parent_dir = current_dir
for _ in range(node.level - 1):
parent_dir = os.path.dirname(parent_dir)
if node.module:
SearchExtensions = ['.pyi', '.py']
sub_path_base = os.path.join(parent_dir, node.module)
found_sub = False
for ext in SearchExtensions:
candidate = sub_path_base + ext
if os.path.isfile(candidate):
actual_module = f"{module_name}.{node.module}" if module_name else node.module
self._RegisterSubModuleSha1(candidate, actual_module, node.module, Gen)
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(actual_module)
for alias in node.names:
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(candidate, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=module_name)
else:
self._LoadModuleDeclarationsFromFile(candidate, Gen, module_name=alias.name, register_module_name=register_module_name, actual_module_name=actual_module, reexport_package=module_name)
has_functions_or_classes = True
found_sub = True
break
# SHA1 map 回退:当文件查找失败时,从 SHA1 映射中查找子模块 stub
if not found_sub:
actual_module = f"{module_name}.{node.module}" if module_name else node.module
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
# 尝试完整模块名和短模块名
target_sha1 = ModuleSha1Map.get(actual_module) or ModuleSha1Map.get(node.module)
if target_sha1:
temp_dir = getattr(Gen, '_temp_dir', None)
if temp_dir:
sha1_pyi = os.path.join(temp_dir, f"{target_sha1}.pyi")
if os.path.isfile(sha1_pyi):
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(actual_module)
for alias in node.names:
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=module_name)
else:
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=alias.name, register_module_name=register_module_name, actual_module_name=actual_module, reexport_package=module_name)
has_functions_or_classes = True
else:
mod_dir = parent_dir
SearchExtensions = ['.pyi', '.py']
# Load __init__.py to make package-level names available
self._LoadPackageInitForRelativeImport(mod_dir, Gen, register_module_name or module_name)
for alias in node.names:
found_alias = False
for ext in SearchExtensions:
sub_path = os.path.join(mod_dir, alias.name + ext)
if os.path.isfile(sub_path):
sub_ModulePath = f"{module_name}.{alias.name}" if module_name else alias.name
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(sub_ModulePath)
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(sub_path, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=module_name)
else:
self._LoadModuleDeclarationsFromFile(sub_path, Gen, module_name=alias.name, register_module_name=register_module_name, reexport_package=module_name)
has_functions_or_classes = True
found_alias = True
break
# SHA1 map 回退
if not found_alias:
sub_ModulePath = f"{module_name}.{alias.name}" if module_name else alias.name
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
target_sha1 = ModuleSha1Map.get(sub_ModulePath) or ModuleSha1Map.get(alias.name)
if target_sha1:
temp_dir = getattr(Gen, '_temp_dir', None)
if temp_dir:
sha1_pyi = os.path.join(temp_dir, f"{target_sha1}.pyi")
if os.path.isfile(sha1_pyi):
if not getattr(self.Trans, '_ImportedModules', None):
self.Trans._ImportedModules = set()
self.Trans._ImportedModules.add(sub_ModulePath)
if alias.name == '*':
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=module_name, register_module_name=register_module_name, reexport_package=module_name)
else:
self._LoadModuleDeclarationsFromFile(sha1_pyi, Gen, module_name=alias.name, register_module_name=register_module_name, reexport_package=module_name)
has_functions_or_classes = True
except Exception as e:
if module_name == 'pyzlib':
node_type = type(node).__name__
name = getattr(node, 'name', getattr(node, 'attr', '?'))
_vlog().warning(f"HandlesImports: 忽略异常 {e}", exc_info=e)
continue
if not has_functions_or_classes and module_name:
self._LoadDeclarationsFromStubLlvm(module_name, Gen)
elif 'includes' in FullModulePath and module_name:
self._LoadDeclarationsFromStubLlvm(module_name, Gen)
except Exception as e:
if module_name:
self._LoadDeclarationsFromStubLlvm(module_name, Gen)
def _EmitExternalGlobalDeclLlvm(self, Node: ast.AnnAssign | ast.Assign, Gen: LlvmGeneratorMixin, module_name: str | None = None, ModulePath: str | None = None) -> None:
"""Emit external global variable declaration from imported module"""
if isinstance(Node, ast.AnnAssign) and isinstance(Node.target, ast.Name):
VarName = Node.target.id
# 不再跳过 _ 开头的全局变量,跨模块引用需要声明它们
var_type = ir.IntType(32)
IsPtr = False
is_string_val = False
TypeInfo = None
if VarName in Gen.module.globals:
existing = Gen.module.globals[VarName]
if isinstance(existing, ir.GlobalVariable):
return
if Node.annotation:
parse_result = ParseListAnnotation(Node.annotation)
if parse_result and not parse_result.is_pointer:
# list[type, count] 数组声明
elem_type_node = parse_result.elem_type_node
count_node = parse_result.count_node
ElemTypeInfo = CTypeInfo.FromNode(elem_type_node, self.Trans.SymbolTable)
if ElemTypeInfo:
elem_type = Gen._ctype_to_llvm(ElemTypeInfo)
if isinstance(elem_type, ir.VoidType):
elem_type = ir.IntType(8)
if isinstance(count_node, ast.Constant) and isinstance(count_node.value, int):
var_type = ir.ArrayType(elem_type, count_node.value)
elif isinstance(count_node, ast.Constant) and count_node.value is None:
var_type = ir.ArrayType(elem_type, 0)
else:
var_type = Gen._ctype_to_llvm(ElemTypeInfo)
IsPtr = False
else:
TypeInfo = CTypeInfo.FromNode(Node.annotation, self.Trans.SymbolTable)
if TypeInfo:
var_type = Gen._ctype_to_llvm(TypeInfo)
IsPtr = TypeInfo.IsPtr
if isinstance(var_type, ir.IdentifiedStructType) and (var_type.elements is None or len(var_type.elements) == 0):
var_type = ir.IntType(32)
if Node.value and ModulePath and IsPtr and isinstance(Node.value, ast.Constant) and isinstance(Node.value.value, str):
str_val = Node.value.value + '\x00'
str_bytes = str_val.encode('utf-8')
arr_type = ir.ArrayType(ir.IntType(8), len(str_bytes))
str_gv_name = VarName + '_str'
if str_gv_name not in Gen.module.globals:
str_gv = ir.GlobalVariable(Gen.module, arr_type, name=str_gv_name)
str_gv.initializer = ir.Constant(arr_type, bytearray(str_bytes))
str_gv.linkage = 'internal'
ptr_type = ir.PointerType(ir.IntType(8))
if VarName not in Gen.module.globals:
ptr_gv = ir.GlobalVariable(Gen.module, ptr_type, name=VarName)
str_gv = Gen.module.globals[str_gv_name]
ptr_gv.initializer = ir.Constant.gep(str_gv, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), 0)])
ptr_gv.linkage = 'internal'
Gen.variables[VarName] = None
return
# 检查是否是 CDefine 常量或 CTypedef 类型别名
IsCDefine = False
IsCTypedef = False
if Node.annotation:
if isinstance(Node.annotation, ast.Attribute):
if getattr(Node.annotation.value, 'id', None) == 't' and Node.annotation.attr == 'CDefine':
IsCDefine = True
elif getattr(Node.annotation.value, 'id', None) == 't' and Node.annotation.attr == 'CTypedef':
IsCTypedef = True
elif isinstance(Node.annotation, ast.Name):
if Node.annotation.id == 'CDefine':
IsCDefine = True
elif Node.annotation.id == 'CTypedef':
IsCTypedef = True
if IsCTypedef:
# 检查符号是否已存在且已正确解析OriginalType 已设置),避免覆盖
existing_sym = self.Trans.SymbolTable.lookup(VarName) if hasattr(self.Trans.SymbolTable, 'lookup') else None
if existing_sym and isinstance(existing_sym, CTypeInfo) and existing_sym.IsTypedef and existing_sym.OriginalType:
return
if Node.value:
ValueTypeInfo = None
if isinstance(Node.value, ast.BinOp) and isinstance(Node.value.op, ast.BitOr):
ValueTypeInfo = self.Trans.TypeMergeHandler.MergeTypes(Node.value)
else:
ValueTypeInfo = self.Trans.TypeMergeHandler.GetCTypeInfo(Node.value)
if isinstance(ValueTypeInfo, CTypeInfo) and (ValueTypeInfo.BaseType or ValueTypeInfo.PtrCount > 0):
ValueTypeInfo.IsTypedef = True
ValueTypeInfo.Name = VarName
self.Trans.SymbolTable.insert(VarName, ValueTypeInfo)
if module_name:
FullName = f"{module_name}.{VarName}"
self.Trans.SymbolTable.insert(FullName, ValueTypeInfo)
return
TTypeInfo = self.Trans.TypeMergeHandler.GetCTypeInfo(Node.annotation)
if TTypeInfo:
TTypeInfo.IsTypedef = True
TTypeInfo.Name = VarName
self.Trans.SymbolTable.insert(VarName, TTypeInfo)
if module_name:
FullName = f"{module_name}.{VarName}"
self.Trans.SymbolTable.insert(FullName, TTypeInfo)
return
if IsCDefine:
# CDefine 宏不生成 LLVM 全局变量,只注册到 SymbolTable
# 引用时会直接替换为立即数
DefineValue = None
if Node.value:
DefineValue = self._ExtractConstValue(Node.value)
# 注册不带前缀的键名(例如 Z_NO_COMPRESSION
self._RegisterCDefineSymbol(VarName, DefineValue, Node.lineno, ModulePath or 'unknown')
# 如果提供了模块名,也注册带模块前缀的键名
if module_name:
FullName = f"{module_name}.{VarName}"
self._RegisterCDefineSymbol(FullName, DefineValue, Node.lineno, ModulePath or 'unknown')
return # 不再继续生成 LLVM 全局变量
# t.State 标记的全局变量只声明不定义external linkage避免 DSO/non-DSO 重定义
if TypeInfo and TypeInfo.IsState:
if VarName not in Gen.module.globals:
gv = ir.GlobalVariable(Gen.module, var_type, name=VarName)
gv.linkage = 'external'
Gen._export_funcs.add(VarName)
return
gv = ir.GlobalVariable(Gen.module, var_type, name=VarName)
# 外部全局变量声明只声明不定义external linkage避免多重定义
gv.linkage = 'external'
elif isinstance(node, ast.Assign) and Node.targets and isinstance(Node.targets[0], ast.Name):
VarName = Node.targets[0].id
if VarName.startswith('_'):
return
if VarName in Gen.module.globals:
return
gv = ir.GlobalVariable(Gen.module, ir.IntType(32), name=VarName)
gv.linkage = 'external'
def _ExtractValue(self, value_node: ast.AST, target_type: ir.Type, var_name: str | None = None) -> ir.Constant | None:
"""Extract constant value from AST node for global variable initializer"""
# 处理负号
if isinstance(value_node, ast.UnaryOp) and isinstance(value_node.op, ast.USub):
if isinstance(value_node.operand, ast.Constant) and isinstance(value_node.operand.value, int):
neg_val = -value_node.operand.value
if isinstance(target_type, ir.IntType):
return ir.Constant(target_type, neg_val)
return ir.Constant(ir.IntType(32), neg_val)
# 处理正号
elif isinstance(value_node, ast.UnaryOp) and isinstance(value_node.op, ast.UAdd):
return self._ExtractValue(value_node.operand, target_type)
if isinstance(value_node, ast.Constant):
if isinstance(value_node.value, bool):
if isinstance(target_type, ir.IntType):
return ir.Constant(target_type, 1 if value_node.value else 0)
return ir.Constant(ir.IntType(32), 1 if value_node.value else 0)
elif isinstance(value_node.value, int):
if isinstance(target_type, ir.IntType):
return ir.Constant(target_type, value_node.value)
return ir.Constant(ir.IntType(32), value_node.value)
elif isinstance(value_node.value, str):
str_val = value_node.value + '\x00'
str_bytes = str_val.encode('utf-8')
arr_type = ir.ArrayType(ir.IntType(8), len(str_bytes))
Gen = self.Trans.LlvmGen
gv_name = var_name + '_str' if var_name else f"str_global_{id(value_node)}"
if gv_name in Gen.module.globals:
gv = Gen.module.globals[gv_name]
else:
gv = ir.GlobalVariable(Gen.module, arr_type, name=gv_name)
gv.initializer = ir.Constant(arr_type, bytearray(str_bytes))
gv.linkage = 'internal'
ptr_type = ir.PointerType(ir.IntType(8))
return ir.Constant(ptr_type, gv.reference)
elif isinstance(value_node.value, bool):
return ir.Constant(ir.IntType(32), 1 if value_node.value else 0)
elif value_node.value is None:
if isinstance(target_type, ir.BaseStructType):
return ir.Constant(target_type, None)
if isinstance(target_type, ir.PointerType):
return ir.Constant(target_type, None)
return ir.Constant(target_type, 0)
elif isinstance(value_node, ast.Name):
if value_node.id == 'None':
if isinstance(target_type, ir.BaseStructType):
return ir.Constant(target_type, None)
if isinstance(target_type, ir.PointerType):
return ir.Constant(target_type, None)
return ir.Constant(target_type, 0)
return None
def _ExtractConstValue(self, value_node: ast.AST) -> object:
"""Extract constant value from AST node for SymbolTable registration
支持常量、一元运算符、二元运算符和常量引用
"""
if isinstance(value_node, ast.Constant):
return value_node.value
elif isinstance(value_node, ast.BinOp):
left = self._ExtractConstValue(value_node.left)
right = self._ExtractConstValue(value_node.right)
if left is None or right is None:
return None
if isinstance(value_node.op, ast.Add):
return left + right
elif isinstance(value_node.op, ast.Sub):
return left - right
elif isinstance(value_node.op, ast.Mult):
return left * right
elif isinstance(value_node.op, ast.Div):
return left // right if isinstance(left, int) and isinstance(right, int) else left / right
elif isinstance(value_node.op, ast.FloorDiv):
return left // right
elif isinstance(value_node.op, ast.Mod):
return left % right
elif isinstance(value_node.op, ast.Pow):
return left ** right
elif isinstance(value_node.op, ast.LShift):
return left << right
elif isinstance(value_node.op, ast.RShift):
return left >> right
elif isinstance(value_node.op, ast.BitOr):
return left | right
elif isinstance(value_node.op, ast.BitXor):
return left ^ right
elif isinstance(value_node.op, ast.BitAnd):
return left & right
elif isinstance(value_node, ast.UnaryOp):
operand = self._ExtractConstValue(value_node.operand)
if operand is None:
return None
if isinstance(value_node.op, ast.USub):
return -operand
elif isinstance(value_node.op, ast.UAdd):
return +operand
elif isinstance(value_node.op, ast.Invert):
return ~operand
elif isinstance(value_node, ast.Name):
if value_node.id in getattr(self.Trans, 'SymbolTable', {}):
info = self.Trans.SymbolTable[value_node.id]
if info.IsDefine and info.DefineValue is not None:
return info.DefineValue
elif isinstance(value_node, ast.Call):
if isinstance(value_node.func, ast.Attribute):
if getattr(value_node.func.value, 'id', None) == 't':
if value_node.args:
return self._ExtractConstValue(value_node.args[0])
elif isinstance(value_node.func, ast.Name):
if value_node.args:
return self._ExtractConstValue(value_node.args[0])
return None
def _ExtractTypedefOriginalType(self, value_node: ast.AST) -> str | None:
if isinstance(value_node, ast.Attribute):
if getattr(value_node.value, 'id', None) == 't':
llvm_str = CTypeRegistry.NameToLLVM(value_node.attr)
if llvm_str:
return llvm_str
elif isinstance(value_node, ast.Name):
llvm_str = CTypeRegistry.NameToLLVM(value_node.id)
if llvm_str:
return llvm_str
elif isinstance(value_node, ast.BinOp) and isinstance(value_node.op, ast.BitOr):
left = self._ExtractTypedefOriginalType(value_node.left)
right = self._ExtractTypedefOriginalType(value_node.right)
parts = []
if left:
parts.append(left)
if right:
parts.append(right)
if parts:
return '|'.join(parts)
elif isinstance(value_node, ast.Call):
if isinstance(value_node.func, ast.Attribute):
if getattr(value_node.func.value, 'id', None) == 't':
llvm_str = CTypeRegistry.NameToLLVM(value_node.func.attr)
if llvm_str:
return llvm_str
elif isinstance(value_node.func, ast.Name):
llvm_str = CTypeRegistry.NameToLLVM(value_node.func.id)
if llvm_str:
return llvm_str
return None
def _RegisterCDefineSymbol(self, FullName: str, DefineValue: object, lineno: int, FilePath: str) -> None:
"""Register CDefine constant to SymbolTable"""
info = CTypeInfo()
info.IsDefine = True
info.DefineValue = DefineValue
info.Lineno = lineno
info.file = FilePath
# 直接添加到 SymbolTable 字典
self.Trans.SymbolTable.insert(FullName, info)
def _check_annotation_for_state(self, annotation: ast.AST) -> bool:
return (AnnotationContainsName(annotation, 'CExport') or
AnnotationContainsName(annotation, 'CExtern') or
AnnotationContainsName(annotation, 'State'))
def _EmitExternalFuncDeclLlvm(self, Node: ast.FunctionDef, Gen: LlvmGeneratorMixin, is_class_method: bool = False, source_module_name: str | None = None, register_module_name: str | None = None, reexport_module_names: list[str] | None = None) -> None:
FuncName = Node.name
if Node.returns and self._check_annotation_for_state(Node.returns):
Gen._export_funcs.add(FuncName) # t.State 标记的外部 C 函数必须保持原始名称,以便链接器解析
if FuncName in Gen.functions:
return
CReturnTypes = []
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)
IsPtr = False
ReturnTypeInfo = None
if Node.returns:
try:
ReturnTypeInfo = self.ResolveAnnotationType(Node.returns)
if ReturnTypeInfo and ReturnTypeInfo.IsStr:
ReturnTypeInfo = CTypeInfo()
ReturnTypeInfo.BaseType = t.CChar()
ReturnTypeInfo.PtrCount = 1
if ReturnTypeInfo is None:
ReturnTypeInfo = CTypeInfo()
ReturnTypeInfo.BaseType = t.CInt()
IsPtr = ReturnTypeInfo.IsPtr
except Exception as e: # 回退:返回类型解析失败时使用默认 CInt
_vlog().warning(f"HandlesImports: 忽略异常 {e}", exc_info=e)
ReturnTypeInfo = CTypeInfo()
ReturnTypeInfo.BaseType = t.CInt()
else:
ReturnTypeInfo = CTypeInfo()
ReturnTypeInfo.BaseType = t.CInt()
ReturnType = Gen._ctype_to_llvm(ReturnTypeInfo)
if (isinstance(ReturnType, ir.VoidType) or (isinstance(ReturnType, ir.PointerType) and isinstance(ReturnType.pointee, ir.VoidType))) and Node.returns:
RetAnnName = None
if isinstance(Node.returns, ast.Name):
RetAnnName = Node.returns.id
elif isinstance(Node.returns, ast.Attribute):
RetAnnName = Node.returns.attr
elif isinstance(Node.returns, ast.BinOp) and isinstance(Node.returns.op, ast.BitOr):
left = Node.returns.left
if isinstance(left, ast.Name):
RetAnnName = left.id
elif isinstance(left, ast.Attribute):
RetAnnName = left.attr
if RetAnnName and RetAnnName in Gen.structs:
ReturnType = ir.PointerType(Gen.structs[RetAnnName])
elif IsPtr or RetAnnName:
ReturnType = ir.PointerType(ir.IntType(8))
ParamTypes = []
is_method = is_class_method or '.__' in FuncName
class_name_for_method = FuncName.split('.')[0] if '.' in FuncName else None
class_is_cpython = class_name_for_method and self.Trans.SymbolTable.is_struct(class_name_for_method)
for i, Arg in enumerate(Node.args.args):
if i == 0 and is_method:
# self parameter of a method should always be a pointer to the struct
ParamType = ir.PointerType(Gen.structs.get(class_name_for_method or FuncName.split('.')[0], ir.IntType(8)))
elif Arg.annotation:
try:
ParamTypeInfo = self.ResolveAnnotationType(Arg.annotation)
if ParamTypeInfo is None:
ParamTypeInfo = CTypeInfo()
ParamTypeInfo.BaseType = t.CInt()
ParamIsPtr = ParamTypeInfo.IsPtr
if i == 0 and is_method and Arg.arg == 'self' and class_is_cpython:
ParamIsPtr = True
ParamTypeInfo.PtrCount = max(ParamTypeInfo.PtrCount, 1)
if ParamTypeInfo.IsStr:
ParamTypeInfo = CTypeInfo()
ParamTypeInfo.BaseType = t.CChar()
ParamTypeInfo.PtrCount = 1
ParamIsPtr = True
# C 语言中数组参数退化为指针list[type, N] → type*
if ParamTypeInfo.ArrayDims:
ParamTypeInfo.ArrayDims = []
ParamTypeInfo.PtrCount = max(ParamTypeInfo.PtrCount, 1)
ParamType = Gen._ctype_to_llvm(ParamTypeInfo)
except Exception as e: # 回退:参数类型解析失败时使用默认 i32
_vlog().warning(f"HandlesImports: 忽略异常 {e}", exc_info=e)
ParamType = ir.IntType(32)
if isinstance(ParamType, ir.VoidType) or (isinstance(ParamType, ir.PointerType) and isinstance(ParamType.pointee, ir.VoidType)):
AnnName = None
if isinstance(Arg.annotation, ast.Name):
AnnName = Arg.annotation.id
elif isinstance(Arg.annotation, ast.Attribute):
AnnName = Arg.annotation.attr
elif isinstance(Arg.annotation, ast.BinOp) and isinstance(Arg.annotation.op, ast.BitOr):
AnnName = ExtractTypeNameFromBinOp(Arg.annotation)
if AnnName and AnnName in Gen.structs:
ParamType = ir.PointerType(Gen.structs[AnnName])
else:
ParamType = ir.PointerType(ir.IntType(8))
else:
ParamType = ir.IntType(32)
ParamTypes.append(ParamType)
if CReturnTypes:
for _ in CReturnTypes:
ParamTypes.append(ir.PointerType(ir.IntType(8)))
IsVariadic = Node.args.vararg is not None
FuncType = ir.FunctionType(ReturnType, ParamTypes, var_arg=IsVariadic)
_is_export = False
if isinstance(ReturnTypeInfo, CTypeInfo) and ReturnTypeInfo.IsState:
_is_export = True
if isinstance(ReturnTypeInfo, CTypeInfo) and ReturnTypeInfo.Storage and isinstance(ReturnTypeInfo.Storage, t.CExport) and not ReturnTypeInfo.IsState:
Gen._export_funcs.add(FuncName)
_is_export = True
if isinstance(ReturnTypeInfo, CTypeInfo) and ReturnTypeInfo.Storage and isinstance(ReturnTypeInfo.Storage, t.CExtern):
Gen._export_funcs.add(FuncName)
_is_export = True
if not _is_export and Node.returns:
_is_export = self._check_annotation_for_state(Node.returns)
MangledName = Gen._mangle_func_name(FuncName, module_name=source_module_name)
func = ir.Function(Gen.module, FuncType, name=MangledName)
Gen.functions[MangledName] = func
Gen.functions[FuncName] = func
func_meta = FuncMeta.NONE
if Node.decorator_list:
for d in Node.decorator_list:
if isinstance(d, ast.Name):
if d.id == 'staticmethod':
func_meta |= FuncMeta.STATIC_METHOD
elif d.id == 'property':
func_meta |= FuncMeta.PROPERTY_GETTER
elif d.id == 'classmethod':
func_meta |= FuncMeta.CLASS_METHOD
elif isinstance(d, ast.Attribute):
if isinstance(d.value, ast.Name):
# 支持 @property.setter 和 @propname.setter 两种形式
if d.value.id == 'property' or d.attr in ('setter', 'getter', 'deleter'):
if d.attr == 'setter':
func_meta |= FuncMeta.PROPERTY_SETTER
elif d.attr == 'getter':
func_meta |= FuncMeta.PROPERTY_GETTER
elif d.attr == 'deleter':
func_meta |= FuncMeta.PROPERTY_DELETER
if not self.Trans.SymbolTable.has(FuncName):
FuncInfo = CTypeInfo()
FuncInfo.Name = FuncName
FuncInfo.IsFunction = True
FuncInfo.MetaList = func_meta
self.Trans.SymbolTable.insert(FuncName, FuncInfo)
else:
existing = self.Trans.SymbolTable[FuncName]
if existing.MetaList == FuncMeta.NONE and func_meta != FuncMeta.NONE:
existing.MetaList = func_meta
# 同时注册带模块前缀的符号,以便 module.func_name 查找能命中
if source_module_name and source_module_name not in ('c', 't'):
FullSymKey = f"{source_module_name}.{FuncName}"
if not self.Trans.SymbolTable.has(FullSymKey):
FullFuncInfo = CTypeInfo()
FullFuncInfo.Name = FullSymKey
FullFuncInfo.IsFunction = True
FullFuncInfo.MetaList = func_meta
self.Trans.SymbolTable.insert(FullSymKey, FullFuncInfo)
# property setter/deleter: 在原始 PropKey不带后缀下注册 MetaList
if FuncMeta.PROPERTY_SETTER in func_meta or FuncMeta.PROPERTY_DELETER in func_meta:
BasePropKey = FuncName.replace('$set', '').replace('$del', '')
base_existing = self.Trans.SymbolTable.lookup(BasePropKey)
if base_existing:
if func_meta != FuncMeta.NONE:
base_existing.MetaList = base_existing.MetaList | func_meta
else:
PropInfo = CTypeInfo()
PropInfo.Name = BasePropKey
PropInfo.IsFunction = True
PropInfo.MetaList = func_meta
self.Trans.SymbolTable.insert(BasePropKey, PropInfo)
if register_module_name and register_module_name != source_module_name:
ReexportName = Gen._mangle_func_name(FuncName, module_name=register_module_name)
Gen.functions[ReexportName] = func
if reexport_module_names:
for reexport_mod in reexport_module_names:
if reexport_mod and reexport_mod != source_module_name:
ReexportName2 = Gen._mangle_func_name(FuncName, module_name=reexport_mod)
if ReexportName2 not in Gen.functions:
Gen.functions[ReexportName2] = func
def _EmitExternalClassDeclLlvm(self, Node: ast.ClassDef, Gen: LlvmGeneratorMixin, module_name: str | None = None, actual_module_name: str | None = None) -> None:
ClassName = Node.name
if hasattr(Node, 'type_params') and Node.type_params:
return
IsCenum = False
IsCpythonObject = False
IsCVTable = False
if Node.bases:
for base in Node.bases:
if getattr(base, 'attr', None):
if base.attr == 'CEnum' or base.attr == 'Enum' or base.attr == 'REnum':
IsCenum = True
break
elif getattr(base, 'id', None):
if base.id == 'CEnum' or base.id == 'Enum' or base.id == 'REnum':
IsCenum = True
break
if not IsCenum and getattr(Node, 'decorator_list', None):
for decorator in Node.decorator_list:
if isinstance(decorator, ast.Attribute):
if getattr(decorator.value, 'id', None) == 't':
if decorator.attr == 'Object':
IsCpythonObject = True
elif decorator.attr == 'CVTable':
IsCVTable = True
elif isinstance(decorator, ast.Call):
if isinstance(decorator.func, ast.Attribute):
if getattr(decorator.func.value, 'id', None) == 't':
if decorator.func.attr == 'Object':
IsCpythonObject = True
elif decorator.func.attr == 'CVTable':
IsCVTable = True
elif isinstance(decorator.func, ast.Name):
if decorator.func.id == 'Object':
IsCpythonObject = True
elif decorator.func.id == 'CVTable':
IsCVTable = True
if ClassName.startswith('__') and not IsCpythonObject and not IsCenum:
return
IsPacked = False
if getattr(Node, 'decorator_list', None):
for decorator in Node.decorator_list:
if isinstance(decorator, ast.Call):
if isinstance(decorator.func, ast.Attribute):
if getattr(decorator.func.value, 'id', None) == 'c' and decorator.func.attr == 'Attribute':
for arg in decorator.args:
if isinstance(arg, ast.Attribute):
if isinstance(arg.value, ast.Attribute):
if getattr(arg.value.value, 'id', None) == 't' and arg.value.attr == 'attr' and arg.attr == 'packed':
IsPacked = True
if IsPacked:
Gen.class_packed.add(ClassName)
if IsCpythonObject:
CInfo = CTypeInfo()
CInfo.Name = ClassName
CInfo.IsCpythonObject = True
CInfo.IsStruct = True
if module_name:
FullName = f"{module_name}.{ClassName}"
self.Trans.SymbolTable.insert(FullName, CInfo)
self.Trans.SymbolTable.insert(ClassName, CInfo)
if IsCenum:
EnumTypeNode = CTypeInfo()
EnumTypeNode.Name = ClassName
EnumTypeNode.BaseType = t.CEnum(ClassName)
EnumTypeNode.IsEnum = True
self.Trans.SymbolTable.insert(ClassName, EnumTypeNode)
for item in Node.body:
VarName = None
if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name):
VarName = item.target.id
elif isinstance(item, ast.Assign):
for target in item.targets:
if isinstance(target, ast.Name):
VarName = target.id
break
if VarName is None:
continue
value = 0
if isinstance(item, ast.Assign) and isinstance(item.value, ast.Constant) and isinstance(item.value.value, int):
value = item.value.value
MemberNode = CTypeInfo()
MemberNode.Name = VarName
MemberNode.BaseType = t.CEnum(ClassName)
MemberNode.value = value
MemberNode.EnumName = ClassName
MemberNode.Lineno = item.lineno
MemberNode.IsEnumMember = True
self.Trans.SymbolTable.insert(VarName, MemberNode)
self.Trans.SymbolTable.insert(f"{ClassName}.{VarName}", MemberNode)
self.Trans.SymbolTable.insert(f"{ClassName}_{VarName}", MemberNode)
return
self._TryLoadStructFromStub(ClassName, Gen)
source_sha1 = None
if actual_module_name and hasattr(Gen, 'ModuleSha1Map') and actual_module_name in Gen.ModuleSha1Map:
source_sha1 = Gen.ModuleSha1Map[actual_module_name]
elif module_name and hasattr(Gen, 'ModuleSha1Map') and module_name in Gen.ModuleSha1Map:
source_sha1 = Gen.ModuleSha1Map[module_name]
# Track which module defines this class (for cross-module name mangling)
if source_sha1:
if not hasattr(Gen, 'class_sha1_map'):
Gen.class_sha1_map = {}
Gen.class_sha1_map[ClassName] = source_sha1
Gen._get_or_create_struct(ClassName, source_sha1=source_sha1, packed=IsPacked)
if ClassName not in Gen.class_members:
Gen.class_members[ClassName] = []
if ClassName not in Gen.class_member_defaults:
Gen.class_member_defaults[ClassName] = {}
if ClassName not in Gen.class_member_signeds:
Gen.class_member_signeds[ClassName] = {}
if ClassName not in Gen.class_member_bitfields:
Gen.class_member_bitfields[ClassName] = {}
if ClassName not in Gen.class_member_byteorders:
Gen.class_member_byteorders[ClassName] = {}
if ClassName not in Gen.class_member_bitoffsets:
Gen.class_member_bitoffsets[ClassName] = {}
has_methods = False
# 解析 __provides__/__requires__/__require_must__ 编译期元数据(跨模块加载)
for item in Node.body:
if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name):
field_name: str = item.target.id
if field_name in ('__provides__', '__requires__', '__require_must__') and item.value:
values: list[str] = []
if isinstance(item.value, ast.List):
for elt in item.value.elts:
if isinstance(elt, ast.Constant) and isinstance(elt.value, str):
values.append(elt.value)
if field_name == '__provides__':
Gen.class_provides[ClassName] = values
elif field_name == '__requires__':
Gen.class_requires[ClassName] = values
elif field_name == '__require_must__':
Gen.class_require_must[ClassName] = values
_existing_member_names = {m[0] for m in Gen.class_members[ClassName]}
for item in Node.body:
# 跳过编译期元数据字段__provides__/__requires__/__require_must__
if isinstance(item, ast.AnnAssign) and isinstance(item.target, ast.Name):
if item.target.id in ('__provides__', '__requires__', '__require_must__'):
continue
parsed = self.CollectAnnAssignMember(item, Gen)
if parsed:
VarName, MemberType, TypeInfo = parsed
if TypeInfo is None:
if VarName not in _existing_member_names:
Gen.class_members[ClassName].append((VarName, MemberType))
_existing_member_names.add(VarName)
Gen.class_member_signeds[ClassName][VarName] = None
Gen.class_member_bitfields[ClassName][VarName] = 0
else:
if TypeInfo.IsBitField:
Gen.class_member_bitfields[ClassName][VarName] = TypeInfo.BitWidth
else:
if VarName not in _existing_member_names:
Gen.class_members[ClassName].append((VarName, MemberType))
_existing_member_names.add(VarName)
Gen.class_member_bitfields[ClassName][VarName] = 0
if TypeInfo.ByteOrder:
Gen.class_member_byteorders[ClassName][VarName] = TypeInfo.ByteOrder
else:
Gen.class_member_byteorders[ClassName][VarName] = ""
if item.value:
const = self._BuildScalarConstant(item.value, MemberType, Gen)
if const:
Gen.class_member_defaults[ClassName][VarName] = const
elif isinstance(item, ast.FunctionDef):
has_methods = True
MethodName = item.name
FuncFullName = f"{ClassName}.__init__" if MethodName == "__init__" else f"{ClassName}.__call__" if MethodName == "__call__" else f"{ClassName}.{MethodName}"
is_item_static = any(isinstance(d, ast.Name) and d.id == 'staticmethod' for d in item.decorator_list)
is_item_property = any(isinstance(d, ast.Name) and d.id == 'property' for d in item.decorator_list)
is_item_classmethod = any(isinstance(d, ast.Name) and d.id == 'classmethod' for d in item.decorator_list)
is_item_prop_setter = any(isinstance(d, ast.Attribute) and d.attr == 'setter' and isinstance(d.value, ast.Name) and d.value.id == 'property' for d in item.decorator_list)
is_item_prop_getter = any(isinstance(d, ast.Attribute) and d.attr == 'getter' and isinstance(d.value, ast.Name) and d.value.id == 'property' for d in item.decorator_list)
is_item_prop_deleter = any(isinstance(d, ast.Attribute) and d.attr == 'deleter' and isinstance(d.value, ast.Name) and d.value.id == 'property' for d in item.decorator_list)
# setter/deleter 使用不同的函数名后缀
DeclFuncName = FuncFullName
if is_item_prop_setter:
DeclFuncName = FuncFullName + '$set'
elif is_item_prop_deleter:
DeclFuncName = FuncFullName + '$del'
if DeclFuncName not in Gen.functions:
FuncDeclNode = ast.FunctionDef(
name=DeclFuncName,
args=item.args,
body=item.body,
decorator_list=item.decorator_list,
returns=item.returns
)
self._EmitExternalFuncDeclLlvm(FuncDeclNode, Gen, is_class_method=not is_item_static and not is_item_classmethod, source_module_name=module_name)
func_meta = FuncMeta.NONE
if is_item_static:
func_meta |= FuncMeta.STATIC_METHOD
if is_item_property:
func_meta |= FuncMeta.PROPERTY_GETTER
if is_item_classmethod:
func_meta |= FuncMeta.CLASS_METHOD
if is_item_prop_setter:
func_meta |= FuncMeta.PROPERTY_SETTER
if is_item_prop_getter:
func_meta |= FuncMeta.PROPERTY_GETTER
if is_item_prop_deleter:
func_meta |= FuncMeta.PROPERTY_DELETER
# setter/deleter 的 SymKey 使用带后缀的函数名
SymKey = DeclFuncName
if not self.Trans.SymbolTable.has(SymKey):
FuncInfo = CTypeInfo()
FuncInfo.Name = SymKey
FuncInfo.IsFunction = True
FuncInfo.MetaList = func_meta
self.Trans.SymbolTable.insert(SymKey, FuncInfo)
else:
existing = self.Trans.SymbolTable[SymKey]
if existing.MetaList == FuncMeta.NONE and func_meta != FuncMeta.NONE:
existing.MetaList = func_meta
# setter/deleter: 同时在原始 PropKey不带后缀下注册 MetaList
if is_item_prop_setter or is_item_prop_deleter:
base_existing = self.Trans.SymbolTable.lookup(FuncFullName)
if base_existing:
if func_meta != FuncMeta.NONE:
base_existing.MetaList = base_existing.MetaList | func_meta
else:
PropInfo = CTypeInfo()
PropInfo.Name = FuncFullName
PropInfo.IsFunction = True
PropInfo.MetaList = func_meta
self.Trans.SymbolTable.insert(FuncFullName, PropInfo)
if has_methods:
if IsCVTable:
Gen._cross_module_vtable_classes.add(ClassName)
Gen.class_vtable.add(ClassName)
NewFuncName = f'{ClassName}.__before_init__'
if not Gen._has_function(NewFuncName) and (IsCpythonObject or IsCVTable):
source_sha1 = None
if module_name and hasattr(Gen, 'ModuleSha1Map') and module_name in Gen.ModuleSha1Map:
source_sha1 = Gen.ModuleSha1Map[module_name]
MangledName = Gen._mangle_name(NewFuncName) if not source_sha1 else f"{source_sha1}.{NewFuncName}"
StructType = Gen.structs.get(ClassName)
if StructType:
StructPtrType = ir.PointerType(StructType)
NewFuncType = ir.FunctionType(ir.VoidType(), [StructPtrType])
NewFunc = ir.Function(Gen.module, NewFuncType, name=MangledName)
Gen.functions[NewFuncName] = NewFunc
bitfields = Gen.class_member_bitfields.get(ClassName, {})
if any(v > 0 for v in bitfields.values()):
current_bit_offset = 0
for name, _ in Gen.class_members.get(ClassName, []):
bw = bitfields.get(name, 0)
Gen.class_member_bitoffsets[ClassName][name] = current_bit_offset
current_bit_offset += bw
def _TryLoadStructFromStub(self, class_name: str, Gen: LlvmGeneratorMixin) -> None:
# 缓存键包含 Gen 实例标识,确保每个模块的 Gen 都能独立加载 struct 定义,
# 同时仍能防止同一 Gen 内的递归重入。原先使用纯 class_name 作为键会导致
# 跨模块 struct 类型对象无法分别 set_body每个模块有独立的 self.structs
cache_key = (id(Gen), class_name)
if cache_key in self._struct_Load_cache:
return
# 先设置缓存防止递归重入,但如果加载失败则清除缓存允许后续重试
self._struct_Load_cache[cache_key] = True
ProjectRoot = self._find_project_root()
temp_dir = os.path.join(ProjectRoot, 'temp')
if not os.path.isdir(temp_dir):
del self._struct_Load_cache[cache_key]
return
if self._stub_cache_dir != temp_dir:
self._stub_cache.clear()
self._stub_cache_dir = temp_dir
for filename in os.listdir(temp_dir):
if not filename.endswith('.stub.ll'):
continue
stub_path = os.path.join(temp_dir, filename)
source_sha1 = filename.replace('.stub.ll', '')
try:
with open(stub_path, 'r', encoding='utf-8') as f:
content = f.read()
for line in content.splitlines():
stripped = line.strip()
match = re.match(r'%"?((?:[a-f0-9]+\.)?(\w+))"?\s*=\s*type\s*(.*)', stripped)
if match:
full_name = match.group(1)
short_name = match.group(2)
struct_body = match.group(3).strip()
self._stub_cache.setdefault(short_name, []).append((full_name, struct_body, source_sha1))
if '.' in full_name:
self._stub_cache.setdefault(full_name, []).append((full_name, struct_body, source_sha1))
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
if class_name not in self._stub_cache:
del self._struct_Load_cache[cache_key]
return
for full_name, struct_body, source_sha1 in self._stub_cache[class_name]:
is_packed = struct_body.startswith('<{') and struct_body.endswith('}>')
if struct_body and struct_body != 'opaque':
elem_types = []
if is_packed:
inner = struct_body[2:-1].strip()
else:
inner = struct_body.strip('{}')
if inner:
for elem in inner.split(','):
elem = elem.strip()
if elem:
et = self._parse_llvm_type(elem, Gen, source_sha1=source_sha1)
elem_types.append(et if et else ir.IntType(32))
try:
st = Gen._get_or_create_struct(class_name, source_sha1=source_sha1, packed=is_packed)
if isinstance(st, ir.IdentifiedStructType) and st.is_opaque:
st.set_body(*elem_types)
if is_packed:
Gen.class_packed.add(class_name)
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
stub_filename = source_sha1 + '.stub.ll'
self._TryLoadClassMembersFromPyi(class_name, stub_filename, Gen)
return
def _TryLoadClassMembersFromPyi(self, class_name: str, stub_filename: str, Gen: LlvmGeneratorMixin) -> None:
ProjectRoot = self._find_project_root()
temp_dir = os.path.join(ProjectRoot, 'temp')
if class_name in Gen.class_members and len(Gen.class_members[class_name]) > 0:
return
pyi_basename = stub_filename.replace('.stub.ll', '.pyi')
pyi_path = os.path.join(temp_dir, pyi_basename)
if self._pyi_cache_dir != temp_dir:
self._pyi_cache.clear()
self._pyi_cache_dir = temp_dir
if pyi_path not in self._pyi_cache:
if not os.path.isfile(pyi_path):
self._pyi_cache[pyi_path] = None
return
try:
with open(pyi_path, 'r', encoding='utf-8') as f:
self._pyi_cache[pyi_path] = ast.parse(f.read())
except Exception as e:
_vlog().warning(f"HandlesImports: 忽略异常 {e}", exc_info=e)
self._pyi_cache[pyi_path] = None
return
pyi_tree = self._pyi_cache[pyi_path]
if pyi_tree is None:
return
source_sha1 = stub_filename.replace('.stub.ll', '')
module_name = None
if hasattr(Gen, 'ModuleSha1Map'):
for mod_name, mod_sha1 in Gen.ModuleSha1Map.items():
if mod_sha1 == source_sha1 and '.' not in mod_name:
module_name = mod_name
break
if module_name is None:
for mod_name, mod_sha1 in Gen.ModuleSha1Map.items():
if mod_sha1 == source_sha1:
module_name = mod_name
break
if module_name is None and hasattr(Gen, 'module_sha1') and Gen.module_sha1 == source_sha1:
module_name = '__self__'
found = False
for node in pyi_tree.body:
if isinstance(node, ast.ClassDef) and node.name == class_name:
self._EmitExternalClassDeclLlvm(node, Gen, module_name=module_name)
found = True
break
if not found:
return
if class_name in Gen.class_members and len(Gen.class_members[class_name]) > 0:
return
return
def _BuildScalarConstant(self, value_node: ast.AST, target_type: ir.Type, Gen: LlvmGeneratorMixin) -> ir.Constant | None:
"""Build a scalar constant value from AST node, delegated to unified method."""
return self.Trans._BuildScalarConstantUnified(value_node, target_type, Gen)
def _LookupStubFuncType(self, func_name: str, Gen: LlvmGeneratorMixin) -> ir.FunctionType | None:
"""从 stub.ll 文件中查找指定函数的 LLVM 类型"""
if not getattr(self, '_stub_func_cache', None):
self._stub_func_cache = {}
ProjectRoot = self._find_project_root()
temp_dir = os.path.join(ProjectRoot, 'temp')
if not os.path.isdir(temp_dir):
return None
if func_name not in self._stub_func_cache:
for filename in os.listdir(temp_dir):
if not filename.endswith('.stub.ll'):
continue
stub_path = os.path.join(temp_dir, filename)
source_sha1 = filename.replace('.stub.ll', '')
try:
with open(stub_path, 'r', encoding='utf-8') as f:
content = f.read()
for line in content.splitlines():
stripped = line.strip()
if not stripped.startswith('declare '):
continue
name_match = stripped.split('@', 1)
if len(name_match) < 2:
continue
stub_func_name = name_match[1].split('(', 1)[0].strip().strip('"')
match = re.match(r'declare\s+(.+?)\s+@("?[\w.]+"?)\s*\((.*)\)$', stripped)
if match:
stub_ret = match.group(1).strip()
stub_params = match.group(3).strip()
self._stub_func_cache[stub_func_name] = (stub_ret, stub_params, '...' in stub_params, source_sha1)
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
if func_name not in self._stub_func_cache:
return None
ret_str, params_str, is_variadic, source_sha1 = self._stub_func_cache[func_name]
ret_type = self._parse_llvm_type(ret_str, Gen, source_sha1=source_sha1, create_structs=False)
if ret_type is None:
ret_type = ir.IntType(32)
param_types = []
if params_str and params_str != 'void':
depth = 0
current = ''
for ch in params_str:
if ch == ',' and depth == 0:
p = self._strip_llvm_param_name(current.strip())
if p and p != '...':
pt = self._parse_llvm_type(p, Gen, source_sha1=source_sha1, create_structs=False)
if pt is not None:
param_types.append(pt)
else:
param_types.append(ir.PointerType(ir.IntType(8)))
current = ''
else:
if ch == '{':
depth += 1
elif ch == '}':
depth -= 1
current += ch
p = self._strip_llvm_param_name(current.strip())
if p and p != '...':
pt = self._parse_llvm_type(p, Gen, source_sha1=source_sha1, create_structs=False)
if pt is not None:
param_types.append(pt)
else:
param_types.append(ir.PointerType(ir.IntType(8)))
return ir.FunctionType(ret_type, param_types, var_arg=is_variadic)
def _strip_llvm_param_name(self, param_str: str) -> str:
"""剥离 LLVM IR 参数字符串中的参数名,仅保留类型部分"""
param_str = param_str.strip()
result = re.sub(r'\s+%[\w.]+\s*$', '', param_str)
return result.strip()
def _parse_simple_llvm_type(self, type_str: str) -> ir.Type | None:
"""解析简单的 LLVM 类型字符串(不创建结构体,避免副作用)"""
type_str = type_str.strip()
type_map = {
'void': ir.VoidType(),
'i1': ir.IntType(1),
'i8': ir.IntType(8),
'i16': ir.IntType(16),
'i32': ir.IntType(32),
'i64': ir.IntType(64),
'float': ir.FloatType(),
'double': ir.DoubleType(),
}
if type_str in type_map:
return type_map[type_str]
if type_str.endswith('*'):
pointee = type_str[:-1].strip()
pointee_type = self._parse_simple_llvm_type(pointee)
if pointee_type:
return ir.PointerType(pointee_type)
return ir.PointerType(ir.IntType(8))
if type_str.startswith('[') and ']' in type_str:
arr_match = re.match(r'\[(\d+)\s*x\s+(.+)\]', type_str)
if arr_match:
size = int(arr_match.group(1))
elem_type = self._parse_simple_llvm_type(arr_match.group(2).strip())
if elem_type:
return ir.ArrayType(elem_type, size)
if type_str.startswith('%'):
return ir.PointerType(ir.IntType(8))
return None
def _LoadDeclarationsFromStubLlvm(self, module_name: str, Gen: LlvmGeneratorMixin) -> None:
"""从 stub.ll 文件加载函数声明和结构体定义到 Gen"""
ProjectRoot = self._find_project_root()
temp_dir = os.path.join(ProjectRoot, 'temp')
if not os.path.isdir(temp_dir):
return
all_stub_files = [f for f in os.listdir(temp_dir) if f.endswith('.stub.ll')]
source_sig_files = getattr(self.Trans, '_source_module_sig_files', {})
target_sha1 = None
if module_name:
for key, val in source_sig_files.items():
if key == module_name or key.endswith('.' + module_name):
target_sha1 = os.path.basename(val).replace('.pyi', '')
break
if not target_sha1:
ModuleSha1Map = getattr(Gen, 'ModuleSha1Map', {})
target_sha1 = ModuleSha1Map.get(module_name)
if target_sha1:
target_stub = f"{target_sha1}.stub.ll"
if target_stub in all_stub_files:
stub_files = [target_stub]
else:
stub_files = []
else:
stub_files = []
needed_sha1s = set()
for filename in stub_files:
stub_path = os.path.join(temp_dir, filename)
try:
with open(stub_path, 'r', encoding='utf-8') as f:
content = f.read()
for line in content.splitlines():
stripped = line.strip()
if stripped.startswith('%') and '=' in stripped and 'type' in stripped:
for m in re.finditer(r'%\"?([a-f0-9]{16})\.', stripped):
needed_sha1s.add(m.group(1))
elif stripped.startswith('declare '):
for m in re.finditer(r'%\"?([a-f0-9]{16})\.', stripped):
needed_sha1s.add(m.group(1))
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
preLoad_files = []
for sha1 in needed_sha1s:
dep_stub = f"{sha1}.stub.ll"
if dep_stub in all_stub_files and dep_stub not in stub_files:
preLoad_files.append(dep_stub)
Load_order = preLoad_files + stub_files
for filename in Load_order:
stub_path = os.path.join(temp_dir, filename)
source_sha1 = filename.replace('.stub.ll', '')
is_preLoad = filename in preLoad_files
try:
with open(stub_path, 'r', encoding='utf-8') as f:
content = f.read()
for line in content.splitlines():
stripped = line.strip()
if stripped.startswith('%') and '=' in stripped and 'type' in stripped:
struct_def_match = re.match(r'%"?((?:[a-f0-9]+\.)?[\w.]+)"?\s*=\s*type\s*(.*)', stripped)
if struct_def_match:
raw_name = struct_def_match.group(1)
if '.' in raw_name:
clean_name = raw_name.split('.', 1)[1]
else:
clean_name = raw_name
struct_body = struct_def_match.group(2).strip()
if struct_body and struct_body != 'opaque':
elem_types = []
is_packed = struct_body.startswith('<{') and struct_body.endswith('}>')
if is_packed:
inner = struct_body[2:-1].strip()
else:
inner = struct_body.strip('{}')
if inner:
for elem in inner.split(','):
elem = elem.strip()
if elem:
elem_type = self._parse_llvm_type(elem, Gen, source_sha1=source_sha1)
if elem_type is not None:
elem_types.append(elem_type)
else:
elem_types.append(ir.IntType(32))
st = Gen._get_or_create_struct(clean_name, source_sha1=source_sha1, packed=is_packed)
if isinstance(st, ir.IdentifiedStructType) and st.elements is None:
try:
st.set_body(*elem_types)
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
continue
if not stripped.startswith('declare '):
continue
if is_preLoad:
continue
func_name_match = stripped.split('@', 1)
if len(func_name_match) < 2:
continue
func_name = func_name_match[1].split('(', 1)[0].strip().strip('"')
if not module_name:
continue
try:
func_type = self._parse_llvm_declare(stripped, Gen)
if func_type is not None:
if func_name in Gen.functions:
existing_func = Gen.functions[func_name]
existing_ftype = getattr(existing_func, 'ftype', None) or existing_func.type.pointee
if str(existing_ftype) != str(func_type):
func = ir.Function(Gen.module, func_type, name=func_name)
Gen.functions[func_name] = func
if '.' in func_name:
short_name = func_name.split('.', 1)[1]
Gen.functions[short_name] = func
else:
func = ir.Function(Gen.module, func_type, name=func_name)
Gen.functions[func_name] = func
if '.' in func_name:
short_name = func_name.split('.', 1)[1]
Gen.functions[short_name] = func
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"处理导入失败: {_e}", "Exception")
def _TryLoadFuncDeclsFromOutputStub(self, struct_name: str, Gen: LlvmGeneratorMixin) -> bool:
"""按需从 output 目录的 stub.ll 加载泛型特化方法声明。
当 temp stub 不包含泛型特化(如 list[str].__iter__
从 output stub由 _split_ll 从完整 IR 生成)中按需加载。
struct_name: 结构体名,如 '9163064cf3eb88f4.list[str]''list[str]'
返回 True 如果找到了并加载了新函数声明。
"""
ProjectRoot = self._find_project_root()
output_dir = os.path.join(ProjectRoot, 'output')
if not os.path.isdir(output_dir):
return False
# 提取短名称和 SHA1 前缀
sha1_prefix: str | None = None
if '.' in struct_name:
parts: list[str] = struct_name.split('.', 1)
if len(parts[0]) == 16 and all(c in '0123456789abcdef' for c in parts[0]):
sha1_prefix = parts[0]
short_name: str = parts[1]
else:
short_name = struct_name
else:
short_name = struct_name
# 确定要搜索的 stub 文件列表
stub_files_to_search: list[str] = []
if sha1_prefix:
target_stub = f'{sha1_prefix}.stub.ll'
target_path = os.path.join(output_dir, target_stub)
if os.path.isfile(target_path):
stub_files_to_search.append(target_stub)
else:
# 没有SHA1前缀搜索所有 output stub
for fname in os.listdir(output_dir):
if fname.endswith('.stub.ll'):
stub_files_to_search.append(fname)
loaded_any: bool = False
for fname in stub_files_to_search:
stub_path = os.path.join(output_dir, fname)
try:
with open(stub_path, 'r', encoding='utf-8') as f:
content = f.read()
for line in content.splitlines():
stripped = line.strip()
if not stripped.startswith('declare '):
continue
# 检查是否是目标结构体的方法
if f'.{short_name}.' not in stripped:
continue
try:
func_type = self._parse_llvm_declare(stripped, Gen)
if func_type is None:
continue
# 提取函数名:最后一个 @ 到第一个 ( 之间的内容
paren_pos: int = stripped.index('(')
at_pos: int = stripped.rindex('@', 0, paren_pos)
func_name: str = stripped[at_pos + 1:paren_pos].strip().strip('"')
if func_name not in Gen.functions:
func = ir.Function(Gen.module, func_type, name=func_name)
Gen.functions[func_name] = func
if '.' in func_name:
short_fn = func_name.split('.', 1)[1]
Gen.functions[short_fn] = func
loaded_any = True
except Exception:
pass
except Exception as _e:
if _config_mode == "strict":
raise
_vlog().warning(f"按需加载 output stub 失败: {_e}", "Exception")
return loaded_any
def _parse_llvm_declare(self, declare_line: str, Gen: LlvmGeneratorMixin) -> ir.FunctionType | None:
"""解析 LLVM declare 语句"""
match = re.match(r'declare\s+(.+?)\s+@("?[^()]+"?)\s*\((.*)\)$', declare_line)
if not match:
return None
ret_type_str = match.group(1).strip()
params_str = match.group(3).strip()
ret_type = self._parse_llvm_type(ret_type_str, Gen)
if ret_type is None:
return None
param_types = []
if params_str and params_str != 'void':
depth = 0
current = ''
for ch in params_str:
if ch == ',' and depth == 0:
p = self._strip_llvm_param_name(current.strip())
if p and p != '...':
pt = self._parse_llvm_type(p, Gen)
if pt is not None:
param_types.append(pt)
current = ''
else:
if ch == '{':
depth += 1
elif ch == '}':
depth -= 1
current += ch
p = self._strip_llvm_param_name(current.strip())
if p and p != '...':
pt = self._parse_llvm_type(p, Gen)
if pt is not None:
param_types.append(pt)
is_variadic = '...' in params_str
return ir.FunctionType(ret_type, param_types, var_arg=is_variadic)
def _parse_llvm_type(self, type_str: str, Gen: LlvmGeneratorMixin, source_sha1: str | None = None, create_structs: bool = True) -> ir.Type | None:
"""解析 LLVM 类型字符串"""
type_str = type_str.strip()
type_map = {
'void': ir.VoidType(),
'i1': ir.IntType(1),
'i8': ir.IntType(8),
'i16': ir.IntType(16),
'i32': ir.IntType(32),
'i64': ir.IntType(64),
'float': ir.FloatType(),
'double': ir.DoubleType(),
}
if type_str in type_map:
return type_map[type_str]
if type_str.startswith('{') and type_str.endswith('}'):
inner = type_str[1:-1].strip()
field_types = []
depth = 0
current = ''
for ch in inner:
if ch == ',' and depth == 0:
field_types.append(current.strip())
current = ''
else:
if ch == '{':
depth += 1
elif ch == '}':
depth -= 1
current += ch
if current.strip():
field_types.append(current.strip())
llvm_fields = []
for ft in field_types:
parsed = self._parse_llvm_type(ft, Gen, source_sha1=source_sha1, create_structs=create_structs)
if parsed is not None:
llvm_fields.append(parsed)
else:
return None
if llvm_fields:
return ir.LiteralStructType(llvm_fields)
return None
if type_str.endswith('*'):
pointee = type_str[:-1].strip()
pointee_type = self._parse_llvm_type(pointee, Gen, source_sha1=source_sha1, create_structs=create_structs)
if pointee_type is not None:
return ir.PointerType(pointee_type)
return ir.PointerType(ir.IntType(8))
if type_str.startswith('[') and ']' in type_str:
arr_match = re.match(r'\[(\d+)\s*x\s+(.+)\]', type_str)
if arr_match:
size = int(arr_match.group(1))
elem_type = self._parse_llvm_type(arr_match.group(2).strip(), Gen, source_sha1=source_sha1, create_structs=create_structs)
if elem_type is not None:
return ir.ArrayType(elem_type, size)
typedef_names = {'typedef', 'CType', 'CVolatile', 'CEnum', 'CUnion', 'CStruct', 'enum', 'REnum', 'renum'}
struct_match = re.match(r'%"?((?:[a-f0-9]+\.)?[\w.\[\]]+)"?', type_str)
if struct_match:
raw_name = struct_match.group(1)
actual_source_sha1 = source_sha1
if '.' in raw_name:
parts = raw_name.split('.', 1)
if len(parts[0]) == 16 and all(c in '0123456789abcdef' for c in parts[0]):
actual_source_sha1 = parts[0]
struct_name = parts[1]
else:
struct_name = raw_name
if struct_name not in typedef_names:
# 泛型类特化(如 list[str]):触发特化创建结构体和方法声明
# 使用 declare_only=True 避免在引用模块中生成方法体(定义由定义模块提供)
if '[' in struct_name and actual_source_sha1:
if struct_name in Gen.structs:
return Gen.structs[struct_name]
gc_base: str = struct_name.split('[')[0]
gc_args_raw: str = struct_name[len(gc_base):]
gc_type_args: list[str] = []
for gc_arg_part in gc_args_raw.split(']['):
gc_arg_part = gc_arg_part.strip('[]')
if gc_arg_part:
gc_type_args.append(gc_arg_part)
if (gc_type_args
and hasattr(self.Trans, 'ClassHandler')
and hasattr(self.Trans.ClassHandler, '_generic_class_templates')
and gc_base in self.Trans.ClassHandler._generic_class_templates):
gc_type_names: list[str] = []
for gc_ta in gc_type_args:
gc_tn: str = gc_ta
if gc_ta == 'int':
gc_tn = 'CInt'
elif gc_ta == 'double':
gc_tn = 'CDouble'
elif gc_ta == 'float':
gc_tn = 'CFloat'
elif gc_ta == 'char':
gc_tn = 'CChar'
elif gc_ta == 'bool':
gc_tn = 'CBool'
gc_type_names.append(gc_tn)
saved_module_sha1: str | None = getattr(Gen, 'module_sha1', None)
saved_trans_sha1: str | None = getattr(self.Trans, '_module_sha1', None)
Gen.module_sha1 = actual_source_sha1
self.Trans._module_sha1 = actual_source_sha1
try:
gc_spec: str | None = self.Trans.ClassHandler._specialize_generic_class(
gc_base, gc_type_args, Gen, type_names=gc_type_names,
declare_only=True)
except Exception:
gc_spec = None
Gen.module_sha1 = saved_module_sha1
self.Trans._module_sha1 = saved_trans_sha1
if gc_spec and gc_spec in Gen.structs:
return Gen.structs[gc_spec]
if create_structs:
return Gen._get_or_create_struct(struct_name, source_sha1=actual_source_sha1)
else:
existing = Gen.structs.get(struct_name)
if existing is not None:
return existing
return None
return ir.IntType(32)
return None
def _strip_param_name(self, param_str: str) -> str:
"""从 LLVM 参数字符串中剥离参数名,只保留类型部分
例如: '%"sha1.struct_name"* %param_name' -> '%"sha1.struct_name"*'
'i8* %path' -> 'i8*'
'i32' -> 'i32'
"""
param_str = param_str.strip()
stripped = re.sub(r'\s+%[\w.]+\s*$', '', param_str)
return stripped.strip()