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

285 lines
14 KiB
Python
Raw Permalink 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, Any
if TYPE_CHECKING:
from lib.core.translator import Translator
from lib.core.LlvmCodeGenerator import LlvmCodeGenerator
import ast
import llvmlite.ir as ir
from lib.core.SymbolUtils import IsTModule
class ExprUtils:
def __init__(self, translator: "Translator") -> None:
self.Trans: Translator = translator
def GetOpSymbol(self, Op: ast.cmpop | ast.operator | ast.unaryoperator) -> str | None:
OpMap: dict[type, str] = {
ast.Add: '+', ast.Sub: '-', ast.Mult: '*', ast.Div: '/',
ast.FloorDiv: '//', ast.Mod: '%', ast.Pow: '**', ast.BitAnd: '&',
ast.BitOr: '|', ast.BitXor: '^', ast.LShift: '<<', ast.RShift: '>>',
ast.MatMult: '@',
}
return OpMap.get(type(Op), None)
def GetUnaryOpSymbol(self, Op: ast.unaryoperator) -> str | None:
OpMap: dict[type, str] = {ast.Invert: '~', ast.Not: 'not', ast.UAdd: '+', ast.USub: '-'}
return OpMap.get(type(Op), None)
def GetComparatorSymbol(self, Op: ast.cmpop) -> str | None:
OpMap: dict[type, str] = {
ast.Eq: '==', ast.NotEq: '!=', ast.Lt: '<', ast.LtE: '<=',
ast.Gt: '>', ast.GtE: '>=', ast.Is: 'is', ast.IsNot: 'is not',
ast.In: 'in', ast.NotIn: 'not in',
}
return OpMap.get(type(Op), None)
def _get_var_class(self, node: ast.AST, Gen: LlvmCodeGenerator) -> str | None:
if isinstance(node, ast.Name):
VarName: str = node.id
return Gen.var_struct_class.get(VarName)
if isinstance(node, ast.Attribute):
AttrName: str = node.attr
# 从父对象的类成员类型推导(如 self._ht → ParsedArgs._ht → HashTable
# 必须先查 BaseClass避免字段名和变量名冲突导致错误分派
# (如 old_args.args 匹配到函数参数 args 的 var_struct_class='AST'
BaseClass: str | None = self._get_var_class(node.value, Gen)
if not BaseClass and isinstance(node.value, ast.Name) and node.value.id == 'self':
BaseClass = getattr(self.Trans, '_CurrentCpythonObjectClass', None)
if BaseClass and BaseClass in Gen.class_members:
for m_name, m_type in Gen.class_members[BaseClass]:
if m_name == AttrName:
if isinstance(m_type, ir.PointerType) and isinstance(m_type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType)):
found = Gen.find_struct_by_pointee(m_type.pointee)
if found:
return found[0]
# 联合类型(如 GSList[Param] | t.CPtr降级为 i8* 时,
# 从 class_member_element_class 查找泛型特化类名
element_class_map = getattr(Gen, 'class_member_element_class', {})
element_class = element_class_map.get(BaseClass, {}).get(AttrName)
if element_class:
return element_class
break
# Fallback: 直接查找 var_struct_class放在 BaseClass 查找之后,
# 避免字段名和变量名冲突导致错误分派)
cls: str | None = Gen.var_struct_class.get(AttrName)
if cls:
return cls
return None
return None
def _try_operator_overLoad(self, ClassName: str, op_name: str, obj_val: ir.Value, Gen: LlvmCodeGenerator, other_val: ir.Value | None = None) -> ir.Value | None:
FullMethodName: str = f'{ClassName}.{op_name}'
if not Gen._has_function(FullMethodName):
FullMethodName = f'{ClassName}.{op_name}__'
if not Gen._has_function(FullMethodName):
# 泛型特化 fallback当 ClassName 是 opaque 基础名(如 'ndarray')且
# 无直接匹配方法时,遍历 Gen.structs 查找以 'ClassName[' 开头的特化版本
# (如 'ndarray[double]')。注解 'numpy.ndarray | t.CPtr' 未指定类型参数时,
# var_struct_class 仅记录裸模板名,但方法只在特化版本中定义。
# 选择策略:优先选 obj_val.type.pointee 实际匹配的特化;否则选第一个有方法的特化。
if '[' not in ClassName:
SpecClass: str | None = None
for CN in Gen.structs.keys():
if CN.startswith(f'{ClassName}[') and Gen._has_function(f'{CN}.{op_name}'):
SpecClass = CN
break
if CN.startswith(f'{ClassName}[') and Gen._has_function(f'{CN}.{op_name}__'):
SpecClass = CN
break
if SpecClass:
FullMethodName = f'{SpecClass}.{op_name}'
if not Gen._has_function(FullMethodName):
FullMethodName = f'{SpecClass}.{op_name}__'
if Gen._has_function(FullMethodName):
func: Any = Gen._get_function(FullMethodName)
call_args: list[Any] = [obj_val]
if other_val is not None:
call_args.append(other_val)
if len(func.ftype.args) == len(call_args):
for i, (expected_type, actual_val) in enumerate(zip(func.ftype.args, call_args)):
if isinstance(expected_type, ir.PointerType) and isinstance(actual_val.type, ir.PointerType):
if expected_type != actual_val.type:
if isinstance(expected_type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType)):
call_args[i] = Gen.builder.bitcast(actual_val, expected_type, name=f"bitcast_arg_{i}")
result: Any = Gen.builder.call(func, call_args, name=f"call_{FullMethodName}")
return result
return None
def _llvm_type_to_detailed_string(self, llvm_type: ir.Type, var_name: str = "") -> dict[str, Any]:
info: dict[str, Any] = {
'name': "unknown",
'size': 0,
'signed': False,
'ptr': False
}
if isinstance(llvm_type, ir.IntType):
info['size'] = llvm_type.width // 8
info['signed'] = True
if llvm_type.width == 8:
info['name'] = "char"
elif llvm_type.width == 16:
info['name'] = "short"
elif llvm_type.width == 32:
info['name'] = "int"
elif llvm_type.width == 64:
info['name'] = "long long"
else:
info['name'] = f"int{llvm_type.width}"
elif isinstance(llvm_type, ir.DoubleType):
info['size'] = 8
info['signed'] = True
info['name'] = "double"
elif isinstance(llvm_type, ir.FloatType):
info['size'] = 4
info['signed'] = True
info['name'] = "float"
elif isinstance(llvm_type, ir.PointerType):
info['size'] = 8
info['signed'] = False
info['ptr'] = True
pointee_info: dict[str, Any] = self._llvm_type_to_detailed_string(llvm_type.pointee)
info['name'] = f"{pointee_info['name']}*"
elif isinstance(llvm_type, (ir.LiteralStructType, ir.IdentifiedStructType)):
size: int = 0
for elem in llvm_type.elements:
elem_info: dict[str, Any] = self._llvm_type_to_detailed_string(elem)
size += elem_info['size']
info['size'] = size
info['signed'] = False
if isinstance(llvm_type, ir.IdentifiedStructType):
info['name'] = llvm_type.name
else:
info['name'] = "struct"
elif isinstance(llvm_type, ir.ArrayType):
elem_info = self._llvm_type_to_detailed_string(llvm_type.element)
info['size'] = elem_info['size'] * llvm_type.count
info['signed'] = elem_info['signed']
info['name'] = f"{elem_info['name']}[]"
elif isinstance(llvm_type, ir.VoidType):
info['size'] = 0
info['signed'] = False
info['name'] = "void"
else:
info['size'] = 0
info['signed'] = False
info['name'] = str(llvm_type)
return info
def _infer_expr_llvm_type_full(self, node: ast.AST) -> ir.Type:
Gen: LlvmCodeGenerator = self.Trans.LlvmGen
if isinstance(node, ast.Constant):
if isinstance(node.value, bool):
return ir.IntType(32)
elif isinstance(node.value, int):
return ir.IntType(32)
elif isinstance(node.value, float):
return ir.DoubleType()
elif isinstance(node.value, str):
if getattr(node, 'kind', None) == 'u':
return ir.PointerType(ir.IntType(16))
return ir.PointerType(ir.IntType(8))
return ir.IntType(32)
elif isinstance(node, ast.Name):
if node.id in Gen.variables:
var_ptr: Any = Gen.variables[node.id]
if isinstance(var_ptr.type, ir.PointerType):
return var_ptr.type.pointee
return var_ptr.type
if node.id in Gen._reg_values:
return Gen._reg_values[node.id].type
return ir.IntType(32)
elif isinstance(node, ast.Attribute):
attr_name: str = node.attr
AttrInfo: Any = self.Trans.SymbolTable.lookup(attr_name)
if AttrInfo and getattr(AttrInfo, 'IsEnumMember', None):
return ir.IntType(32)
return ir.IntType(32)
elif isinstance(node, ast.BinOp):
left_type: ir.Type = self._infer_expr_llvm_type_full(node.left)
right_type: ir.Type = self._infer_expr_llvm_type_full(node.right)
if isinstance(left_type, (ir.FloatType, ir.DoubleType)) or isinstance(right_type, (ir.FloatType, ir.DoubleType)):
return ir.DoubleType()
if isinstance(left_type, ir.PointerType) or isinstance(right_type, ir.PointerType):
return ir.PointerType(ir.IntType(8))
return ir.IntType(32)
elif isinstance(node, ast.UnaryOp):
return self._infer_expr_llvm_type_full(node.operand)
elif isinstance(node, ast.Subscript):
val_type: ir.Type = self._infer_expr_llvm_type_full(node.value)
if isinstance(val_type, ir.PointerType):
if isinstance(val_type.pointee, ir.IntType) and val_type.pointee.width == 8:
return ir.IntType(8)
if isinstance(val_type.pointee, ir.ArrayType):
return ir.PointerType(val_type.pointee.element)
return val_type.pointee
if isinstance(val_type, ir.ArrayType):
return val_type.element
if isinstance(val_type, ir.IntType):
return ir.IntType(8)
return ir.IntType(32)
elif isinstance(node, ast.Call):
if isinstance(node.func, ast.Name):
func_name: str = node.func.id
if Gen._has_function(func_name):
fn: Any = Gen._get_function(func_name)
return fn.function_type.return_type
if func_name == 'len':
return ir.IntType(32)
if func_name == 'sizeof':
return ir.IntType(32)
if func_name == 'ord':
return ir.IntType(32)
if func_name == 'chr':
return ir.IntType(8)
if func_name == 'int':
return ir.IntType(32)
if func_name == 'float' or func_name == 'double':
return ir.DoubleType()
elif isinstance(node.func, ast.Attribute):
if isinstance(node.func.value, ast.Name) and node.func.value.id == 'c':
if node.func.attr == 'Deref':
return ir.IntType(64)
if isinstance(node.func.value, ast.Name) and IsTModule(node.func.value.id, self.Trans.SymbolTable):
if node.func.attr == 'CChar':
return ir.PointerType(ir.IntType(8))
return ir.IntType(32)
elif isinstance(node, ast.Compare):
return ir.IntType(1)
elif isinstance(node, ast.BoolOp):
return ir.IntType(1)
elif isinstance(node, ast.IfExp):
return self._infer_expr_llvm_type_full(node.body)
elif isinstance(node, ast.Attribute):
return ir.PointerType(ir.IntType(8))
return ir.IntType(32)
def _collect_names(self, node: ast.AST, bound: set[str], free: set[str] | None = None) -> None:
"""收集 AST 节点中的绑定变量和自由变量
Args:
node: AST 节点
bound: 集合收集绑定变量ast.Store 上下文)
free: 集合收集自由变量ast.Load 上下文),可为 None
"""
if free is None:
free = set()
if isinstance(node, ast.Name):
if isinstance(node.ctx, ast.Store):
bound.add(node.id)
elif isinstance(node.ctx, ast.Load):
free.add(node.id)
elif isinstance(node, (ast.Lambda, ast.FunctionDef, ast.AsyncFunctionDef)):
# 这些节点有自己的作用域,不递归进入
# 但需要收集参数名作为 bound
if hasattr(node, 'args') and node.args:
for arg in node.args.args:
bound.add(arg.arg)
elif getattr(node, '_fields', None):
for field in node._fields:
value: Any = getattr(node, field, None)
if isinstance(value, list):
for item in value:
if isinstance(item, (ast.stmt, ast.expr)):
self._collect_names(item, bound, free)
elif isinstance(value, (ast.stmt, ast.expr)):
self._collect_names(value, bound, free)