snapshot before regression test

This commit is contained in:
t
2026-07-18 19:25:40 +08:00
commit 796222a300
2295 changed files with 206453 additions and 0 deletions

View File

@@ -0,0 +1,266 @@
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 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)