Files
TransPyC/lib/core/ConstEvaluator.py

633 lines
26 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.
"""统一常量表达式求值器
整合了原先分散在 5 处的常量求值逻辑:
- SymbolTable._eval_const_expr
- CTypeInfo.TryEvalConstExpr
- ConstEvalMixin._try_eval_const_expr
- HandlesAssign._eval_const_expr
- HandlesIf._eval_const_expr
ConstEvaluator 提供三个级别的求值:
1. eval_basic(node) — 纯 AST 常量 + 算术(无外部依赖)
2. eval_with_symtab(node) — + SymbolTable 中的 define 查找
3. eval_full(node, ctx) — + Gen._define_constants / _all_define_constants / 平台宏
"""
from __future__ import annotations
import ast
from typing import Any, Callable, Optional, TYPE_CHECKING
import llvmlite.ir as ir
from lib.core.Handles.HandlesBase import CTypeInfo
# 治本修复:导入 t 模块和 CType 基类,用于类型构造调用求值(如 t.CUnsignedLong(-11)
from lib.includes import t as _t_module
from lib.includes.t import CType as _CType_base
if TYPE_CHECKING:
from lib.core.LlvmCodeGenerator import LlvmCodeGenerator
class ConstEvaluator:
"""统一常量表达式求值器"""
# ==================================================================
# Level 2: + SymbolTable define 查找
# ==================================================================
@staticmethod
def eval_with_symtab(node: ast.AST, symtab: Any) -> Optional[Any]:
"""求值常量表达式,支持从 SymbolTable 查找 define 值。
支持: Constant, BinOp, UnaryOp, Name(define), Attribute(define), Call(类型构造)
"""
if isinstance(node, ast.Constant):
return node.value
if isinstance(node, ast.Name):
if hasattr(symtab, '__contains__') and node.id in symtab:
info: CTypeInfo = symtab[node.id]
if info.IsDefine and isinstance(info.DefineValue, int):
return info.DefineValue
return None
if isinstance(node, ast.Attribute):
# 先尝试 t.CSizeT().Size 模式(类型构造调用 + .Size 属性)
size_val = ConstEvaluator._eval_ctype_size_attr(node)
if size_val is not None:
return size_val
return ConstEvaluator._eval_attribute_define(node, symtab)
# 治本修复:支持 t.CUnsignedLong(-11) 等类型构造调用
if isinstance(node, ast.Call):
return ConstEvaluator._eval_ctype_ctor_call(node, symtab)
return ConstEvaluator._eval_arith(node, lambda n: ConstEvaluator.eval_with_symtab(n, symtab))
@staticmethod
def _eval_ctype_ctor_call(node: ast.Call, symtab: Any) -> Optional[Any]:
"""求值类型构造调用,如 t.CUnsignedLong(-11)、CInt(42)。
在无 ctxEvalContext的场景下通过 symtab 和 t 模块直接求值。
按类型的 Size 和 IsSigned 应用位掩码。
"""
func: ast.expr = node.func
func_name: str | None = None
module_name: str | None = None
if isinstance(func, ast.Attribute):
if isinstance(func.value, ast.Name):
module_name = func.value.id
func_name = func.attr
elif isinstance(func, ast.Name):
func_name = func.id
ctype_cls: Any = None
if module_name == 't' and func_name:
ctype_cls = getattr(_t_module, func_name, None)
elif func_name and not module_name:
ctype_cls = getattr(_t_module, func_name, None)
if ctype_cls is None and symtab:
t_type_syms = getattr(symtab, '_t_type_symbols', {})
cls_candidate = t_type_syms.get(func_name)
if cls_candidate is not None:
ctype_cls = cls_candidate
if ctype_cls is None or not (isinstance(ctype_cls, type) and issubclass(ctype_cls, _CType_base)):
return None
if not node.args:
return None
arg_val: Any = ConstEvaluator.eval_with_symtab(node.args[0], symtab)
if arg_val is None:
return None
size: int | None = None
is_signed: bool | None = None
try:
tmp_inst = ctype_cls()
size = getattr(tmp_inst, 'Size', None)
is_signed = getattr(tmp_inst, 'IsSigned', None)
except Exception:
pass
if isinstance(arg_val, int) and size and size > 0:
mask: int = (1 << size) - 1
if is_signed:
masked: int = arg_val & mask
if masked >= (1 << (size - 1)):
masked -= (1 << size)
return masked
else:
return arg_val & mask
return arg_val
# ==================================================================
# Level 3: + Gen._define_constants / _all_define_constants / 平台宏
# ==================================================================
@staticmethod
def eval_full(node: ast.AST, ctx: 'EvalContext') -> Optional[Any]:
"""完整求值,支持 define_constants、平台宏、BoolOp、Compare。
ctx: EvalContext 实例,封装了 Gen、SymbolTable 等上下文。
"""
# BoolOp (HandlesIf 逻辑)
if isinstance(node, ast.BoolOp):
if isinstance(node.op, ast.And):
for val in node.values:
result: Any = ConstEvaluator.eval_full(val, ctx)
if result is None:
return None
if not result:
return 0
return 1
elif isinstance(node.op, ast.Or):
for val in node.values:
result: Any = ConstEvaluator.eval_full(val, ctx)
if result is None:
return None
if result:
return 1
return 0
return None
# Compare (HandlesIf 逻辑)
if isinstance(node, ast.Compare):
left: Any = ConstEvaluator.eval_full(node.left, ctx)
if left is None:
return None
op: ast.cmpop
comparator: ast.AST
for op, comparator in zip(node.ops, node.comparators):
right: Any = ConstEvaluator.eval_full(comparator, ctx)
if right is None:
return None
result: bool | None = ConstEvaluator._eval_compare_op(op, left, right)
if result is None:
return None
if not result:
return 0
left = right
return 1
# Constant — 直接返回原始值bool→0/1 转换由 HandlesIf 自行处理
if isinstance(node, ast.Constant):
return node.value
# Name — 查 define_constants / SymbolTable / 平台宏
if isinstance(node, ast.Name):
return ctx.lookup_name(node.id)
# Attribute — 查 define_constants / SymbolTable / 平台宏
if isinstance(node, ast.Attribute):
# 先尝试 t.CSizeT().Size 模式(类型构造调用 + .Size 属性)
size_val = ConstEvaluator._eval_ctype_size_attr(node)
if size_val is not None:
return size_val
return ctx.lookup_attribute(node)
# UnaryOp
if isinstance(node, ast.UnaryOp):
operand: Any = ConstEvaluator.eval_full(node.operand, ctx)
if operand is None:
return None
if isinstance(node.op, ast.Not):
return 1 if not operand else 0
if isinstance(node.op, ast.USub):
return -operand
if isinstance(node.op, ast.UAdd):
return +operand
if isinstance(node.op, ast.Invert):
return ~operand
return None
# Call — 编译时函数调用(如 ctraits.isptr(x)
if isinstance(node, ast.Call):
return ConstEvaluator._eval_compile_time_call(node, ctx)
# BinOp
return ConstEvaluator._eval_arith(node, lambda n: ConstEvaluator.eval_full(n, ctx))
# ==================================================================
# 内部辅助
# ==================================================================
@staticmethod
def _eval_compile_time_call(node: ast.Call, ctx: 'EvalContext') -> Optional[Any]:
"""处理编译时函数调用,如 ctraits.isptr(x)、t.CUnsignedLong(-11)"""
func: ast.expr = node.func
func_name: str | None = None
module_name: str | None = None
if isinstance(func, ast.Attribute):
if isinstance(func.value, ast.Name):
module_name = func.value.id
func_name = func.attr
elif isinstance(func, ast.Name):
func_name = func.id
# ctraits.isptr(x) — 带模块前缀
# isptr(x) — 通过 from ctraits import isptr 导入
is_isptr: bool = False
if module_name == 'ctraits' and func_name == 'isptr':
is_isptr = True
elif func_name == 'isptr' and not module_name:
# 检查是否通过 from ctraits import isptr 导入
t_c_imported = getattr(ctx.translator, '_t_c_imported_names', {}) if ctx.translator else {}
if func_name in t_c_imported:
src_module, _ = t_c_imported[func_name]
if src_module == 'ctraits':
is_isptr = True
# 也检查符号表中是否有 ctraits.isptr
if not is_isptr and ctx.symtab:
sym = ctx.symtab.lookup('ctraits.isptr') if hasattr(ctx.symtab, 'lookup') else None
if sym:
is_isptr = True
if is_isptr and node.args and ctx.translator:
try:
arg_node: ast.expr = node.args[0]
# 局部变量:优先通过 LLVM 类型判断GetCTypeInfo 无法解析局部变量名)
if ctx.Gen and isinstance(arg_node, ast.Name):
var_name: str = arg_node.id
variables: dict = getattr(ctx.Gen, 'variables', {})
if var_name in variables:
var_val: Any = variables[var_name]
if hasattr(var_val, 'type') and isinstance(var_val.type, ir.PointerType):
pointee: ir.Type = var_val.type.pointee
if isinstance(pointee, ir.PointerType):
return 1
return 0
# 类型名 / typedef通过 TypeMergeHandler 获取 CTypeInfo
type_info = ctx.translator.TypeMergeHandler.GetCTypeInfo(arg_node)
if type_info and (type_info.IsPtr or type_info.PtrCount > 0):
return 1
if type_info and not (type_info.IsPtr or type_info.PtrCount > 0):
return 0
return None
except Exception:
return None
# 治本修复:支持 t.CUnsignedLong(-11) 等类型构造调用
# 当 t.CUnsignedLong(value) / CUnsignedInt(value) / CInt(value) 等被用作常量表达式时,
# 求值参数并返回(按类型的无符号/有符号语义应用位掩码)。
# 这修复了 stdint.py 中 INFINITE = t.CUnsignedLong(-1) 等常量无法被求值的问题。
is_ctype_ctor: bool = False
ctype_cls: Any = None
if module_name == 't' and func_name:
ctype_cls = getattr(_t_module, func_name, None)
if ctype_cls is not None and isinstance(ctype_cls, type) and issubclass(ctype_cls, _CType_base):
is_ctype_ctor = True
elif func_name and not module_name:
# from t import CUnsignedLong 形式
ctype_cls = getattr(_t_module, func_name, None)
if ctype_cls is not None and isinstance(ctype_cls, type) and issubclass(ctype_cls, _CType_base):
is_ctype_ctor = True
# 也检查 _t_type_symbols符号表中的 t 模块类型符号)
if not is_ctype_ctor and ctx.symtab:
t_type_syms = getattr(ctx.symtab, '_t_type_symbols', {})
cls_candidate = t_type_syms.get(func_name)
if cls_candidate is not None and isinstance(cls_candidate, type) and issubclass(cls_candidate, _CType_base):
ctype_cls = cls_candidate
is_ctype_ctor = True
if is_ctype_ctor and node.args:
try:
arg_node: ast.expr = node.args[0]
arg_val: Any = ConstEvaluator.eval_full(arg_node, ctx)
if arg_val is None:
return None
# 按类型的 Size 和 IsSigned 应用位掩码
size: int | None = getattr(ctype_cls, '_Size', None)
is_signed: bool | None = getattr(ctype_cls, '_IsSigned', None)
# 尝试创建临时实例获取 Size/IsSigned
try:
tmp_inst = ctype_cls()
size = getattr(tmp_inst, 'Size', size)
is_signed = getattr(tmp_inst, 'IsSigned', is_signed)
except Exception:
pass
if isinstance(arg_val, int) and size and size > 0:
mask: int = (1 << size) - 1
if is_signed:
# 有符号类型:解释为补码
masked: int = arg_val & mask
if masked >= (1 << (size - 1)):
masked -= (1 << size)
return masked
else:
# 无符号类型:直接掩码
return arg_val & mask
return arg_val
except Exception:
return None
return None
@staticmethod
def _eval_arith(node: ast.AST, recurse: Callable[[ast.AST], Any]) -> Optional[Any]:
"""通用算术求值BinOp + UnaryOp"""
if isinstance(node, ast.BinOp):
left: Any = recurse(node.left)
right: Any = recurse(node.right)
if left is None or right is None:
return None
return ConstEvaluator._apply_binop(node.op, left, right)
if isinstance(node, ast.UnaryOp):
operand: Any = recurse(node.operand)
if operand is None:
return None
if isinstance(node.op, ast.USub):
return -operand
if isinstance(node.op, ast.UAdd):
return +operand
if isinstance(node.op, ast.Invert):
return ~operand
return None
return None
@staticmethod
def _apply_binop(op: ast.operator, left: Any, right: Any) -> Optional[Any]:
"""应用二元运算符"""
try:
if isinstance(op, ast.Add):
return left + right
elif isinstance(op, ast.Sub):
return left - right
elif isinstance(op, ast.Mult):
return left * right
elif isinstance(op, ast.Div):
return left // right if isinstance(left, int) and isinstance(right, int) else left / right
elif isinstance(op, ast.FloorDiv):
if isinstance(right, (int, float)) and right == 0:
return None
return left // right
elif isinstance(op, ast.Mod):
if isinstance(right, (int, float)) and right == 0:
return None
return left % right
elif isinstance(op, ast.Pow):
return left ** right
elif isinstance(op, ast.LShift):
return left << right
elif isinstance(op, ast.RShift):
return left >> right
elif isinstance(op, ast.BitOr):
return left | right
elif isinstance(op, ast.BitXor):
return left ^ right
elif isinstance(op, ast.BitAnd):
return left & right
except Exception:
return None
return None
@staticmethod
def _eval_compare_op(op: ast.cmpop, left: Any, right: Any) -> Optional[bool]:
"""应用比较运算符"""
if isinstance(op, ast.Eq):
return left == right
elif isinstance(op, ast.NotEq):
return left != right
elif isinstance(op, ast.Lt):
return left < right
elif isinstance(op, ast.LtE):
return left <= right
elif isinstance(op, ast.Gt):
return left > right
elif isinstance(op, ast.GtE):
return left >= right
return None
@staticmethod
def _eval_ctype_size_attr(node: ast.Attribute) -> Optional[int]:
"""求值 t.CSizeT().Size 等类型构造调用 + .Size 属性模式。
匹配 AST 形态: Attribute(attr='Size', value=Call(func=Attribute(value=Name(id='t'), attr='CSizeT')))
或直接 Call(func=Name(id='CSizeT'))from t import CSizeT 后的使用方式)。
返回类型大小(位数),如 CSizeT().Size 返回 6464 位平台)。
"""
if node.attr != 'Size':
return None
if not isinstance(node.value, ast.Call):
return None
call_node: ast.Call = node.value
if call_node.args or call_node.keywords:
return None
func: ast.expr = call_node.func
type_name: str | None = None
if isinstance(func, ast.Attribute) and isinstance(func.value, ast.Name) and func.value.id == 't':
type_name = func.attr
elif isinstance(func, ast.Name):
type_name = func.id
if not type_name:
return None
# 延迟导入避免循环依赖
try:
from lib.includes.t import CTypeRegistry
except Exception:
return None
cls = CTypeRegistry.GetClassByName(type_name)
if cls is None:
return None
try:
instance = cls()
size_val = getattr(instance, 'Size', None)
if isinstance(size_val, int):
return size_val
except Exception:
pass
return None
@staticmethod
def _eval_attribute_define(node: ast.Attribute, symtab: Any) -> Optional[Any]:
"""从 SymbolTable 查找 Attribute 形式的 define 值"""
parts: list[str] = []
current: ast.AST = node
while isinstance(current, ast.Attribute):
parts.append(current.attr)
current = current.value
if isinstance(current, ast.Name):
parts.append(current.id)
parts.reverse()
attr_name: str = parts[-1]
possible_keys: list[str] = [attr_name, '.'.join(parts)]
for key in possible_keys:
if hasattr(symtab, '__contains__') and key in symtab:
info: CTypeInfo = symtab[key]
if info.IsDefine and isinstance(info.DefineValue, int):
return info.DefineValue
# 模糊匹配
if len(parts) >= 2:
mod_key: str
mod_info: CTypeInfo
for mod_key, mod_info in (symtab.items() if hasattr(symtab, 'items') else []):
if mod_info.IsDefine and isinstance(mod_info.DefineValue, int):
if mod_key.endswith('.' + attr_name) or mod_key == attr_name:
return mod_info.DefineValue
return None
@staticmethod
def _intify(val: Any) -> Any:
"""float → int 转换(仅当 float 值恰好是整数时)"""
if isinstance(val, float) and val == int(val):
return int(val)
return val
class EvalContext:
"""常量求值上下文,封装 Gen / SymbolTable / 平台宏查找逻辑。
由调用方构造并传入 ConstEvaluator.eval_full()。
"""
def __init__(self, Gen: LlvmCodeGenerator | None = None, symtab: Any = None, platform_macros: dict | None = None, translator: Any = None):
self.Gen: LlvmCodeGenerator | None = Gen
self.symtab: Any = symtab
self._platform_macros: dict | None = platform_macros
self.translator: Any = translator
def lookup_name(self, name: str) -> Optional[Any]:
"""查找 Name 节点的值define_constants → SymbolTable → 平台宏"""
# 1. Gen._define_constants
if self.Gen and name in self.Gen._define_constants:
val: Any = self.Gen._define_constants[name]
if isinstance(val, (int, float)):
return ConstEvaluator._intify(val)
# 2. Gen._DefineConstants (大写版本)
if self.Gen and hasattr(self.Gen, '_DefineConstants') and name in self.Gen._DefineConstants:
val: Any = self.Gen._DefineConstants[name]
if isinstance(val, (int, float)):
return ConstEvaluator._intify(val)
# 3. SymbolTable
if self.symtab and name in self.symtab:
info: CTypeInfo = self.symtab[name]
val: Any = info.value
if isinstance(val, (int, float)):
return ConstEvaluator._intify(val)
if info.IsDefine and isinstance(info.DefineValue, (int, float)):
return ConstEvaluator._intify(info.DefineValue)
# 4. 平台宏
macros: dict = self._get_platform_macros()
if name in macros:
return macros[name]
return None
def lookup_attribute(self, node: ast.Attribute) -> Optional[Any]:
"""查找 Attribute 节点的值"""
module_name: str | None = node.value.id if isinstance(node.value, ast.Name) else None
attr_name: str = node.attr
if module_name and attr_name:
combined_key: str = f"{module_name}.{attr_name}"
# 1. Gen._define_constants
if self.Gen and combined_key in self.Gen._define_constants:
val: Any = self.Gen._define_constants[combined_key]
if isinstance(val, (int, float)):
return ConstEvaluator._intify(val)
# 2. SymbolTable (精确匹配 + 后缀匹配)
if self.symtab:
key: str
info: CTypeInfo
for key, info in self.symtab.items():
if key.endswith(f".{combined_key}") or key == combined_key:
if info.IsDefine and isinstance(info.DefineValue, (int, float)):
return ConstEvaluator._intify(info.DefineValue)
# 3. _all_define_constants
all_dc: dict = self._get_all_define_constants()
if combined_key in all_dc:
val: Any = all_dc[combined_key]
if isinstance(val, (int, float)):
return ConstEvaluator._intify(val)
if attr_name in all_dc:
val: Any = all_dc[attr_name]
if isinstance(val, (int, float)):
return ConstEvaluator._intify(val)
dc_key: str
dc_val: Any
for dc_key, dc_val in all_dc.items():
if dc_key == attr_name or dc_key.endswith(f".{attr_name}"):
if isinstance(dc_val, (int, float)):
return ConstEvaluator._intify(dc_val)
# 4. SymbolTable 短名查找
if self.symtab and attr_name in self.symtab:
info: CTypeInfo = self.symtab[attr_name]
if info.IsDefine and isinstance(info.DefineValue, (int, float)):
return ConstEvaluator._intify(info.DefineValue)
# 5. Gen._define_constants 短名
if self.Gen and attr_name in self.Gen._define_constants:
val: Any = self.Gen._define_constants[attr_name]
if isinstance(val, (int, float)):
return ConstEvaluator._intify(val)
# 6. 平台宏
macros: dict = self._get_platform_macros()
full_key: str | None = self._get_attr_full_name(node)
if full_key:
if full_key in macros:
return macros[full_key]
short_name: str = full_key.split('.')[-1] if '.' in full_key else full_key
if short_name in macros:
return macros[short_name]
return None
def _get_all_define_constants(self) -> dict:
"""获取 _all_define_constants"""
if self.Gen and hasattr(self.Gen, '_all_define_constants'):
return self.Gen._all_define_constants
return {}
def _get_platform_macros(self) -> dict:
"""获取平台宏"""
if self._platform_macros is not None:
return self._platform_macros
if not self.Gen:
return {}
macros: dict[str, int] = {}
pi: dict = self.Gen._platform_info
ptr_size: int = self.Gen.ptr_size
if pi.get('is_windows'):
macros['_WIN32'] = 1
if not pi.get('is_32bit'):
macros['_WIN64'] = 1
macros['WIN64'] = 1
macros['WIN32'] = 1
if pi.get('is_linux'):
macros['__linux__'] = 1
if pi.get('is_macos'):
macros['__APPLE__'] = 1
macros['__MACH__'] = 1
if pi.get('is_x86_64'):
macros['__x86_64__'] = 1
macros['__x86_64'] = 1
macros['_M_X64'] = 1
elif pi.get('is_arm64'):
macros['__aarch64__'] = 1
macros['_M_ARM64'] = 1
if not pi.get('is_32bit') and pi.get('is_linux'):
macros['__LP64__'] = 1
elif pi.get('is_32bit'):
macros['_ILP32'] = 1
macros['SIZEOF_VOID_P'] = ptr_size
macros['__SIZEOF_POINTER__'] = ptr_size
macros['NDEBUG'] = 0
return macros
@staticmethod
def _get_attr_full_name(node: ast.Attribute) -> Optional[str]:
"""从 Attribute 节点提取完整名称 (e.g. 'module.name')"""
parts: list[str] = []
current: ast.AST = node
while isinstance(current, ast.Attribute):
parts.append(current.attr)
current = current.value
if isinstance(current, ast.Name):
parts.append(current.id)
parts.reverse()
return '.'.join(parts) if parts else None