Files
TransPyC/lib/core/SymbolTable.py

508 lines
21 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.Handles.HandlesBase import CTypeInfo, CTypeHelper
from lib.includes import t
from typing import Any, Dict, Iterator
import ast
# 从 SymbolUtils 重导出,保持向后兼容
from lib.core.SymbolUtils import IsTModuleType as _IsTModuleType
from lib.core.SymbolUtils import AnnotationContainsTType as _AnnotationContainsTType
from lib.core.SymbolUtils import CheckAnnotationHasCInline as _CheckAnnotationHasCInline
class SymbolTable:
def __init__(self, translator: "Translator"):
self.translator = translator
self._symbols: Dict[str, CTypeInfo] = {}
from lib.core.DiagnosticCollector import DiagnosticCollector
self.diagnostics = DiagnosticCollector()
def clear(self):
self._symbols.clear()
def get(self, name: str, default=None) -> CTypeInfo:
return self._symbols.get(name, default)
def set(self, name: str, value: Any):
if isinstance(value, CTypeInfo):
self._symbols[name] = value
else:
# 兼容旧代码传入 dict 的情况
info = self._CTypeInfoFromDict(name, value)
self._symbols[name] = info
def update(self, symbols: Dict[str, Any]):
for name, value in symbols.items():
self.set(name, value)
def keys(self) -> Iterator[str]:
return iter(self._symbols.keys())
def values(self) -> Iterator[CTypeInfo]:
return iter(self._symbols.values())
def items(self) -> Iterator[tuple[str, CTypeInfo]]:
return iter(self._symbols.items())
def pop(self, name: str, *args) -> CTypeInfo:
return self._symbols.pop(name, *args)
def __getitem__(self, name: str) -> CTypeInfo:
if name in self._symbols:
return self._symbols[name]
raise KeyError(name)
def __setitem__(self, name: str, value: Any):
self.set(name, value)
def __delitem__(self, name: str):
del self._symbols[name]
def __contains__(self, name: str) -> bool:
return name in self._symbols
def __len__(self) -> int:
return len(self._symbols)
def __iter__(self) -> Iterator[str]:
return iter(self._symbols)
def ToDict(self) -> Dict[str, Any]:
"""序列化为字典格式"""
result = {}
for name, info in self._symbols.items():
entry = {'name': name}
if info.BaseType:
entry['BaseType'] = str(info.BaseType)
if info.PtrCount:
entry['PtrCount'] = info.PtrCount
if info.IsTypedef:
entry['IsTypedef'] = True
if info.IsStruct:
entry['IsStruct'] = True
if info.IsEnum:
entry['IsEnum'] = True
if info.IsUnion:
entry['IsUnion'] = True
if info.IsFunction:
entry['IsFunction'] = True
if info.IsVariable:
entry['IsVariable'] = True
if info.IsDefine:
entry['IsDefine'] = True
if info.OriginalType:
entry['OriginalType'] = str(info.OriginalType)
if info.Lineno:
entry['lineno'] = info.Lineno
if info.file:
entry['file'] = info.file
if info.Members:
entry['members'] = info.Members
extra = info._sm._extra
if extra:
entry.update(extra)
result[name] = entry
return result
def FromDict(self, symbols: Dict[str, Any]):
"""从字典格式反序列化"""
self._symbols.clear()
for name, attrs in symbols.items():
if isinstance(attrs, dict):
info = self._CTypeInfoFromDict(name, attrs)
self._symbols[name] = info
elif isinstance(attrs, CTypeInfo):
self._symbols[name] = attrs
def _CTypeInfoFromDict(self, name: str, attrs: dict) -> CTypeInfo:
"""从旧 dict 格式创建 CTypeInfo兼容旧代码"""
info = CTypeInfo()
info.Name = name
node_type = attrs.get('type', '')
if node_type == 'struct' or attrs.get('IsStruct'):
info.IsStruct = True
elif node_type == 'enum' or attrs.get('IsEnum'):
info.IsEnum = True
elif node_type == 'union' or attrs.get('IsUnion'):
info.IsUnion = True
elif node_type == 'typedef' or attrs.get('IsTypedef'):
info.IsTypedef = True
elif node_type == 'function' or attrs.get('IsFunction'):
info.IsFunction = True
elif node_type == 'variable' or attrs.get('IsVariable'):
info.IsVariable = True
elif node_type == 'define' or attrs.get('IsDefine'):
info.IsDefine = True
elif node_type == 'enum_member' or attrs.get('IsEnumMember'):
info.IsEnumMember = True
if 'PtrCount' in attrs:
info.PtrCount = attrs['PtrCount']
if 'OriginalType' in attrs:
info.OriginalType = attrs['OriginalType']
if 'lineno' in attrs:
info.Lineno = attrs['lineno']
if 'file' in attrs:
info.file = attrs['file']
if 'members' in attrs:
info.Members = attrs['members']
if 'IsPtr' in attrs:
info.IsPtr = attrs['IsPtr']
if 'dims' in attrs:
info.ArrayDims = attrs['dims']
if 'IsCpythonObject' in attrs:
info.IsCpythonObject = attrs['IsCpythonObject']
if 'IsAnonymous' in attrs:
info.IsAnonymous = attrs['IsAnonymous']
if 'IsPacked' in attrs:
info.IsPacked = attrs['IsPacked']
if 'EnumName' in attrs:
info.EnumName = attrs['EnumName']
if 'DefineValue' in attrs:
info.DefineValue = attrs['DefineValue']
# 其他属性存入 _extra
skip_keys = {'type', 'PtrCount', 'OriginalType', 'lineno', 'file',
'members', 'IsPtr', 'dims', 'IsCpythonObject', 'IsAnonymous',
'IsPacked', 'EnumName', 'DefineValue',
'IsStruct', 'IsEnum', 'IsUnion', 'IsTypedef',
'IsFunction', 'IsVariable', 'IsDefine', 'IsEnumMember'}
for key, value in attrs.items():
if key not in skip_keys:
info.set(key, value)
return info
# ==================================================================
# 模块符号加载:四阶段管线
# ==================================================================
def LoadModuleSymbols(self, FullModulePath: str, asname: str, lineno: int = 0) -> list[str]:
return self._LoadSymbolsFromFile(FullModulePath, asname, lineno)
def _LoadSymbolsFromFile(self, FilePath: str, namespace_prefix: str | list[str], lineno: int = 0) -> list[str]:
from lib.core.SymbolExtractor import ASTSymbolExtractor
from lib.core.SymbolTypeResolver import TypeResolver
from lib.core.SymbolInserter import SymbolInserter
from lib.core.SymbolReexporter import PackageReexporter
try:
# Pass 1: 从 AST 提取原始符号信息
extractor = ASTSymbolExtractor(self)
module_symbols = extractor.extract(FilePath)
# Pass 2: 解析类型信息typedef、函数签名、define 值)
resolver = TypeResolver(self)
resolver.resolve(module_symbols)
# Pass 3: 插入符号到主命名空间(无前缀)
inserter = SymbolInserter(self)
loaded = inserter.insert(module_symbols)
# Pass 4: 重新导出到命名空间前缀下
prefixes = [namespace_prefix] if isinstance(namespace_prefix, str) else namespace_prefix
reexporter = PackageReexporter(self)
loaded.extend(reexporter.reexport(module_symbols, prefixes, lineno))
return loaded
except Exception as e:
import traceback
print(traceback.format_exc())
self.diagnostics.error(FilePath, lineno, f"加载模块失败: {e}")
return []
# ==================================================================
# 类型解析辅助方法(供 TypeResolver 通过 SymbolTable 引用调用)
# ==================================================================
def _GetLLVMTypeStr(self, node):
if node is None:
self.diagnostics.warn('', 0, "类型注解为 None回退到 i32")
return 'i32'
if isinstance(node, ast.Name):
if node.id == 'str':
return 'i8*'
elif node.id == 'int':
return 'i32'
elif node.id == 'bool':
return 'i8'
elif node.id == 'float':
return 'double'
elif node.id == 'None':
return 'void'
elif node.id in ('UINT8PTR', 'INT8PTR', 'BYTEPTR'):
return 'i8*'
elif node.id in ('UINT16PTR', 'INT16PTR'):
return 'i16*'
elif node.id in ('UINT32PTR', 'INT32PTR'):
return 'i32*'
elif node.id in ('UINT64PTR', 'INT64PTR'):
return 'i64*'
else:
if node.id in self:
entry = self[node.id]
if entry and entry.IsTypedef and entry.OriginalType:
if isinstance(entry.OriginalType, CTypeInfo) and entry.OriginalType.IsFuncPtr:
return 'i8*'
elif isinstance(entry.OriginalType, CTypeInfo) and entry.OriginalType.BaseType:
return entry.OriginalType.ToString()
elif isinstance(entry.OriginalType, str):
return entry.OriginalType
self.diagnostics.warn('', 0, f"无法解析 Name 类型 '{node.id}',回退到 i32")
return 'i32'
elif isinstance(node, ast.Attribute):
attr_name = node.attr if hasattr(node, 'attr') else ''
from lib.includes.t import CTypeRegistry
llvm_str = CTypeRegistry.NameToLLVM(attr_name)
if llvm_str is not None:
return llvm_str
if attr_name in ('CState', 'CDefine', 'CTypedef', 'CExtern', 'CStatic', 'CConst', 'State'):
return ''
if attr_name in ('CCharPtr', 'CIntPtr', 'CVoidPtr', 'CArrayPtr'):
return 'i8*'
if attr_name == 'CVoidPtr':
return 'i8*'
if attr_name and attr_name[0].isupper() and attr_name not in CTypeRegistry._name_to_class:
return f'%struct.{attr_name}*'
if attr_name and attr_name not in CTypeRegistry._name_to_class:
return f'%struct.{attr_name}*'
self.diagnostics.warn('', 0, f"无法解析 Attribute 类型 '{attr_name}',回退到 i32")
return 'i32'
elif isinstance(node, ast.Subscript):
if isinstance(node.value, ast.Name) and node.value.id == 'tuple':
slice_node = node.slice
elem_types = []
if isinstance(slice_node, ast.Tuple):
for elt in slice_node.elts:
elem_types.append(self._GetLLVMTypeStr(elt))
else:
elem_types.append(self._GetLLVMTypeStr(slice_node))
if elem_types:
return '{ ' + ', '.join(elem_types) + ' }'
if isinstance(node.value, ast.Name) and node.value.id == 'list':
slice_node = node.slice
if isinstance(slice_node, ast.Tuple) and len(slice_node.elts) >= 1:
elem_type = self._GetLLVMTypeStr(slice_node.elts[0])
else:
elem_type = self._GetLLVMTypeStr(slice_node)
return elem_type + '*'
self.diagnostics.warn('', 0, f"无法解析 Subscript 类型,回退到 i32")
return 'i32'
elif isinstance(node, ast.BinOp) and isinstance(node.op, ast.BitOr):
left = self._GetLLVMTypeStr(node.left)
right = self._GetLLVMTypeStr(node.right)
if right and right.endswith('*'):
return right
if left and left.endswith('*'):
return left
if left and left not in ('i32', 'void', 'i64'):
return left
if right and right not in ('i32', 'void', 'i64'):
return right
return left if left else right
elif isinstance(node, ast.Constant):
if isinstance(node.value, bool):
return 'i8'
return 'i32'
elif isinstance(node, ast.Call):
if isinstance(node.func, ast.Attribute):
if hasattr(node.func.value, 'id') and node.func.value.id == 't':
return self._GetLLVMTypeStr(ast.Attribute(value=ast.Name(id='t'), attr=node.func.attr))
self.diagnostics.warn('', 0, "无法解析 Call 类型,回退到 i32")
return 'i32'
self.diagnostics.warn('', 0, f"无法解析类型节点 {type(node).__name__},回退到 i32")
return 'i32'
def _GetFuncRetTypeStr(self, returns_node):
if returns_node is None:
return 'i32'
actual = returns_node
if isinstance(returns_node, ast.BinOp) and isinstance(returns_node.op, ast.BitOr):
left_str = self._GetLLVMTypeStr(returns_node.left)
right_str = self._GetLLVMTypeStr(returns_node.right)
if right_str and right_str.endswith('*'):
return right_str
if left_str and left_str.endswith('*'):
return left_str
if left_str and left_str not in ('i32', 'void', 'i64'):
return left_str
if right_str and right_str not in ('i32', 'void', 'i64'):
return right_str
return left_str
result = self._GetLLVMTypeStr(actual)
return result if result else 'i32'
@staticmethod
def _CheckAnnotationHasCInline(annotation_node) -> bool:
return _CheckAnnotationHasCInline(annotation_node)
def _GetFuncParamTypeStr(self, annotation_node):
if annotation_node is None:
return 'i8*'
actual = annotation_node
if isinstance(annotation_node, ast.BinOp) and isinstance(annotation_node.op, ast.BitOr):
left_str = self._GetLLVMTypeStr(annotation_node.left)
right_str = self._GetLLVMTypeStr(annotation_node.right)
if right_str and right_str.endswith('*'):
return right_str
if left_str and left_str.endswith('*'):
return left_str
if right_str and right_str not in ('i32', 'void', 'i64'):
return right_str
if left_str and left_str not in ('i32', 'void', 'i64'):
return left_str
if left_str and left_str.endswith('*'):
return left_str
return right_str if right_str else left_str
result = self._GetLLVMTypeStr(actual)
return result if result else 'i8*'
def _ResolveTypedefValueType(self, node):
if node is None:
return ''
if isinstance(node, ast.Name):
if node.id in self:
Entry = self[node.id]
if Entry.IsTypedef and Entry.OriginalType:
return Entry.OriginalType
return ''
if isinstance(node, ast.Attribute):
if isinstance(node.value, ast.Name) and node.value.id == 't':
CName = CTypeHelper.GetCName(node.attr)
if CName and CName == '*':
return 'ptr'
if CName:
return CName
return ''
if isinstance(node, ast.BinOp) and isinstance(node.op, ast.BitOr):
left_str = self._ResolveTypedefValueType(node.left)
right_str = self._ResolveTypedefValueType(node.right)
if right_str == 'ptr' or (right_str and right_str.endswith('*')):
base = left_str if left_str and left_str != 'ptr' else ''
return base + ' *' if base else 'void *'
if left_str == 'ptr' or (left_str and left_str.endswith('*')):
base = right_str if right_str and right_str != 'ptr' else ''
return base + ' *' if base else 'void *'
if right_str and right_str not in ('ptr', 'void', 'i32', 'i64'):
return right_str
if left_str and left_str not in ('ptr', 'void', 'i32', 'i64'):
return left_str
return left_str if left_str else right_str
return ''
# ==================================================================
# 符号插入方法(供 SymbolInserter / PackageReexporter 调用)
# ==================================================================
def _InsertClassSymbol(self, FullName: str, TypeKind: str, lineno: int, FilePath: str, members: dict, IsCpythonObject: bool, IsPacked: bool = False):
info = CTypeInfo()
info.Lineno = lineno
info.file = FilePath
info.Members = members if members else {}
if TypeKind == 'struct':
info.IsStruct = True
elif TypeKind == 'union':
info.IsUnion = True
elif TypeKind == 'enum':
info.IsEnum = True
elif TypeKind == 'exception':
info.IsStruct = True
info.IsExceptionClass = True
if IsCpythonObject:
info.IsCpythonObject = True
if IsPacked:
info.IsPacked = True
self._symbols[FullName] = info
def _InsertEnumMemberSymbol(self, FullName: str, EnumName: str, lineno: int, FilePath: str):
info = CTypeInfo()
info.IsEnumMember = True
info.EnumName = EnumName
info.Lineno = lineno
info.file = FilePath
self._symbols[FullName] = info
def _InsertTypedefSymbol(self, FullName: str, OriginalType_kind: str | None, OriginalClass, lineno: int, FilePath: str, members: dict | None = None):
info = CTypeInfo()
info.IsTypedef = True
info.Lineno = lineno
info.file = FilePath
if members:
info.Members = members
if isinstance(OriginalClass, CTypeInfo):
info.OriginalType = OriginalClass
elif OriginalType_kind == 'typedef' and OriginalClass and isinstance(OriginalClass, str) and ('*' in OriginalClass):
OriginalType = OriginalClass
if OriginalType == 'CVoid *':
OriginalType = 'void *'
info.OriginalType = OriginalType
elif OriginalType_kind == 'typedef' and OriginalClass == 'void':
info.OriginalType = 'void *'
elif OriginalType_kind == 'typedef' and OriginalClass:
info.OriginalType = OriginalClass
elif OriginalType_kind and OriginalClass and isinstance(OriginalClass, str):
info.OriginalType = f'{OriginalType_kind} {OriginalClass}'
else:
info.OriginalType = OriginalClass
self._symbols[FullName] = info
def _InsertModuleSymbol(self, name: str, lineno: int, FilePath: str):
info = CTypeInfo()
info.IsModuleAlias = True
info.Lineno = lineno
info.file = FilePath
self._symbols[name] = info
def _InsertFuncSymbol(self, FullName: str, RetType, ParamTypes: list, lineno: int, FilePath: str, IsVariadic: bool = False, IsInline: bool = False):
info = CTypeInfo()
info.IsFunction = True
if isinstance(RetType, str):
info.FuncPtrReturn = CTypeInfo.FromTypeName(RetType) if RetType else CTypeInfo.VoidTypeInfo()
elif isinstance(RetType, CTypeInfo):
info.FuncPtrReturn = RetType
else:
info.FuncPtrReturn = CTypeInfo.VoidTypeInfo()
info.FuncPtrParams = [(f'arg{i}', pt) for i, pt in enumerate(ParamTypes)]
info.IsVariadic = IsVariadic
info.IsInline = IsInline
if IsInline:
info.Storage = t.CInline()
info.Lineno = lineno
info.file = FilePath
self._symbols[FullName] = info
def _InsertDefineSymbol(self, FullName: str, DefineValue, lineno: int, FilePath: str):
info = CTypeInfo()
info.IsDefine = True
info.DefineValue = DefineValue
info.Lineno = lineno
info.file = FilePath
self._symbols[FullName] = info
def _eval_const_expr(self, node):
"""计算常量表达式的值(委托到 ConstEvaluator"""
from lib.core.ConstEvaluator import ConstEvaluator
return ConstEvaluator.eval_with_symtab(node, self)
def _InsertAnonymousSymbol(self, FullName: str, IsUnion: bool, members: dict, lineno: int, FilePath: str):
info = CTypeInfo()
if IsUnion:
info.IsUnion = True
else:
info.IsStruct = True
info.IsAnonymous = True
info.Members = members if members else {}
info.Lineno = lineno
info.file = FilePath
self._symbols[FullName] = info