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,249 @@
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
from lib.core.Handles.HandlesBase import BaseHandle
import ast
import llvmlite.ir as ir
class ExprLambdaHandle(BaseHandle):
def _HandleLambdaLlvm(self, Node: ast.Lambda) -> ir.Value:
Gen: LlvmCodeGenerator = self.Trans.LlvmGen
if not getattr(Gen, '_lambda_counter', None):
Gen._lambda_counter = 0
lambda_id: int = Gen._lambda_counter
Gen._lambda_counter += 1
fn_name: str = f"lambda_fn_{lambda_id}"
captured_vars: list[tuple[str, Any]] = self._get_lambda_captured_vars(Node)
env_types: list[ir.Type] = []
env_var_names: list[str] = []
for var_name, var_val in captured_vars:
env_types.append(var_val.type)
env_var_names.append(var_name)
env_struct_type: ir.LiteralStructType = ir.LiteralStructType(env_types) if env_types else ir.LiteralStructType([])
ret_type: ir.Type = self.Trans.ExprUtils._infer_expr_llvm_type_full(Node.body)
param_types: list[ir.Type] = [ir.PointerType(env_struct_type)]
if Node.args:
for arg in Node.args.args:
param_types.append(ir.IntType(32))
fn_type: ir.FunctionType = ir.FunctionType(ret_type, param_types)
fn: ir.Function = ir.Function(Gen.module, fn_type, name=fn_name)
Gen.functions[fn_name] = fn
entry_block: ir.Block = fn.append_basic_block(name=f"{fn_name}_entry")
prev_builder: Any = Gen.builder
prev_func: Any = Gen.func
prev_vars: dict[str, Any] = dict(Gen.variables)
Gen.builder = ir.IRBuilder(entry_block)
Gen.func = fn
Gen.variables = {}
for i, arg in enumerate(fn.args):
arg.name = f"arg_{i}" if i > 0 else "env"
if env_types:
for idx, var_name in enumerate(env_var_names):
env_ptr: Any = fn.args[0]
zero: ir.Constant = ir.Constant(ir.IntType(32), 0)
member_ptr: Any = Gen.builder.gep(env_ptr, [zero, ir.Constant(ir.IntType(32), idx)], name=f"env_{var_name}")
Loaded: Any = Gen._load(member_ptr, name=f"Load_env_{var_name}")
temp_alloca: Any = Gen._alloca(Loaded.type, name=f"temp_{var_name}")
Gen._store(Loaded, temp_alloca)
Gen.variables[var_name] = temp_alloca
if Node.args:
for i, arg in enumerate(Node.args.args):
Gen.variables[arg.arg] = Gen._alloca(ir.IntType(32), name=arg.arg)
Gen._store(fn.args[i + 1], Gen.variables[arg.arg])
body_val: Any = self.HandleExprLlvm(Node.body)
if body_val:
if isinstance(body_val.type, ir.PointerType) and not isinstance(ret_type, ir.PointerType):
body_val = Gen._load(body_val, name="Load_lambda_ret")
if body_val.type != ret_type:
if isinstance(body_val.type, ir.IntType) and isinstance(ret_type, ir.IntType):
if body_val.type.width < ret_type.width:
body_val = Gen.builder.zext(body_val, ret_type, name="zext_lambda_ret")
elif body_val.type.width > ret_type.width:
body_val = Gen.builder.trunc(body_val, ret_type, name="trunc_lambda_ret")
elif isinstance(body_val.type, ir.PointerType) and isinstance(ret_type, ir.PointerType):
body_val = Gen.builder.bitcast(body_val, ret_type, name="bitcast_lambda_ret")
Gen.builder.ret(body_val)
else:
if isinstance(ret_type, ir.PointerType):
Gen.builder.ret(ir.Constant(ret_type, None))
elif isinstance(ret_type, ir.IntType):
Gen.builder.ret(ir.Constant(ret_type, 0))
elif isinstance(ret_type, (ir.FloatType, ir.DoubleType)):
Gen.builder.ret(ir.Constant(ret_type, 0.0))
elif isinstance(ret_type, ir.BaseStructType):
if ret_type.elements:
zero_val: ir.Constant = ir.Constant(ret_type, [
ir.Constant(et, None) if isinstance(et, (ir.PointerType, ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType))
else ir.Constant(et, 0) if isinstance(et, (ir.IntType, ir.FloatType, ir.DoubleType))
else ir.Constant(et, ir.Undefined) for et in ret_type.elements])
else:
zero_val = ir.Constant(ret_type, None)
Gen.builder.ret(zero_val)
else:
Gen.builder.ret(ir.Constant(ret_type, 0))
Gen.builder = prev_builder
Gen.func = prev_func
Gen.variables = prev_vars
closure_struct_type: ir.LiteralStructType = ir.LiteralStructType([
ir.PointerType(ir.IntType(8)),
ir.PointerType(fn_type)
])
closure_ptr: Any = Gen._alloca(closure_struct_type, name=f"closure_{lambda_id}")
env_ptr: Any = Gen._alloca(env_struct_type, name=f"env_{lambda_id}")
if env_types:
for idx, (var_name, var_val) in enumerate(captured_vars):
zero = ir.Constant(ir.IntType(32), 0)
member_ptr = Gen.builder.gep(env_ptr, [zero, ir.Constant(ir.IntType(32), idx)], name=f"store_env_{var_name}")
Gen._store(var_val, member_ptr)
zero = ir.Constant(ir.IntType(32), 0)
env_field_ptr: Any = Gen.builder.gep(closure_ptr, [zero, zero], name="closure_env_ptr")
env_i8_ptr: Any = Gen.builder.bitcast(env_ptr, ir.PointerType(ir.IntType(8)), name="env_i8_ptr")
Gen._store(env_i8_ptr, env_field_ptr)
fn_field_ptr: Any = Gen.builder.gep(closure_ptr, [zero, ir.Constant(ir.IntType(32), 1)], name="closure_fn_ptr")
Gen._store(fn, fn_field_ptr)
return closure_ptr
def _get_lambda_captured_vars(self, Node: ast.Lambda) -> list[tuple[str, ir.Value]]:
Gen: LlvmCodeGenerator = self.Trans.LlvmGen
captured: list[tuple[str, Any]] = []
free_vars: list[str] = self._get_lambda_free_vars(Node)
for var_name in free_vars:
if var_name in Gen.variables:
var_ptr: Any = Gen.variables[var_name]
if isinstance(var_ptr, ir.AllocaInstr) or isinstance(var_ptr, ir.GlobalVariable):
Loaded: Any = Gen._load(var_ptr, name=f"capture_{var_name}")
captured.append((var_name, Loaded))
return captured
def _get_lambda_free_vars(self, Node: ast.Lambda) -> list[str]:
bound_vars: set[str] = set()
free_vars: set[str] = set()
# 收集 lambda body 中的 bound 和 free 变量
self.Trans.ExprUtils._collect_names(Node.body, bound_vars, free_vars)
# lambda 参数也是 bound 的
if Node.args:
for arg in Node.args.args:
bound_vars.add(arg.arg)
# free 变量 = 在 body 中引用(Load)但不是参数也不是 body 内赋值的变量
result: list[str] = list(free_vars - bound_vars)
return result
def _HandleIfExpLlvm(self, Node: ast.IfExp) -> ir.Value | None:
Gen: LlvmCodeGenerator = self.Trans.LlvmGen
TestVal: Any = self.HandleExprLlvm(Node.test)
if not TestVal:
return None
if not isinstance(TestVal.type, ir.IntType) or TestVal.type.width != 1:
TestVal = Gen.builder.icmp_signed('!=', TestVal, Gen._ZeroConst(TestVal.type), name="ifexp_cond")
ThenBB: ir.Block = Gen.func.append_basic_block(name="ifexp.then")
ElseBB: ir.Block = Gen.func.append_basic_block(name="ifexp.else")
MergeBB: ir.Block = Gen.func.append_basic_block(name="ifexp.end")
Gen.builder.cbranch(TestVal, ThenBB, ElseBB)
# 先在 ThenBB 中生成 BodyVal
Gen.builder.position_at_start(ThenBB)
BodyVal: Any = self.HandleExprLlvm(Node.body)
if not BodyVal:
BodyVal = ir.Constant(ir.IntType(32), 0)
ThenEndBB: ir.Block = Gen.builder.block
# 在 ElseBB 中生成 OrelseVal
Gen.builder.position_at_start(ElseBB)
OrelseVal: Any = self.HandleExprLlvm(Node.orelse)
if not OrelseVal:
OrelseVal = ir.Constant(ir.IntType(32), 0)
ElseEndBB: ir.Block = Gen.builder.block
# 确定 phi 节点的目标类型
result_type: ir.Type = BodyVal.type
if BodyVal.type != OrelseVal.type:
if isinstance(BodyVal.type, ir.IntType) and isinstance(OrelseVal.type, ir.IntType):
result_type = BodyVal.type if BodyVal.type.width >= OrelseVal.type.width else OrelseVal.type
elif isinstance(BodyVal.type, ir.PointerType) and isinstance(OrelseVal.type, ir.PointerType):
# 指针类型,使用 BodyVal 的类型
result_type = BodyVal.type
elif isinstance(BodyVal.type, ir.IntType) and isinstance(OrelseVal.type, ir.PointerType):
# BodyVal 是整数OrelseVal 是指针(字符串字面量)
# 如果指针指向 i8将指针转换为 i8加载字符
if isinstance(OrelseVal.type.pointee, ir.IntType) and OrelseVal.type.pointee.width == 8:
result_type = ir.IntType(8)
else:
result_type = OrelseVal.type
elif isinstance(BodyVal.type, ir.PointerType) and isinstance(OrelseVal.type, ir.IntType):
# BodyVal 是指针字符串字面量OrelseVal 是整数
# 如果指针指向 i8将指针转换为 i8加载字符
if isinstance(BodyVal.type.pointee, ir.IntType) and BodyVal.type.pointee.width == 8:
result_type = ir.IntType(8)
else:
result_type = BodyVal.type
else:
result_type = OrelseVal.type
# 在 ThenBB 末尾添加类型转换(如果需要)和分支
if not ThenEndBB.is_terminated:
Gen.builder.position_at_end(ThenEndBB)
if BodyVal.type != result_type:
if isinstance(BodyVal.type, ir.IntType) and isinstance(result_type, ir.IntType):
if BodyVal.type.width < result_type.width:
BodyVal = Gen.builder.zext(BodyVal, result_type, name="ifexp_body_zext")
else:
BodyVal = Gen.builder.trunc(BodyVal, result_type, name="ifexp_body_trunc")
elif isinstance(BodyVal.type, ir.IntType) and isinstance(result_type, ir.PointerType):
# 如果 result_type 是 i8*,说明 Else 分支是字符串字面量
# 不应该将整数转换为指针,而应该保持整数类型
# 这里将 result_type 改为 i8
result_type = ir.IntType(8)
# 重新处理 Else 分支的类型转换
elif isinstance(BodyVal.type, ir.PointerType) and isinstance(result_type, ir.IntType):
# 指针转整数
BodyVal = Gen.builder.ptrtoint(BodyVal, result_type, name="ifexp_body_ptr2int")
Gen.builder.branch(MergeBB)
# 在 ElseBB 末尾添加类型转换(如果需要)和分支
if not ElseEndBB.is_terminated:
Gen.builder.position_at_end(ElseEndBB)
if OrelseVal.type != result_type:
if isinstance(OrelseVal.type, ir.IntType) and isinstance(result_type, ir.IntType):
if OrelseVal.type.width < result_type.width:
OrelseVal = Gen.builder.zext(OrelseVal, result_type, name="ifexp_orelse_zext")
else:
OrelseVal = Gen.builder.trunc(OrelseVal, result_type, name="ifexp_orelse_trunc")
elif isinstance(OrelseVal.type, ir.IntType) and isinstance(result_type, ir.PointerType):
# 如果 result_type 是 i8*,说明 BodyVal 分支是字符串字面量
# 不应该将整数转换为指针,而应该保持整数类型
# 这里将 result_type 改为 i8
result_type = ir.IntType(8)
# 重新处理 BodyVal 分支的类型转换(已经在上面处理过了)
elif isinstance(OrelseVal.type, ir.PointerType) and isinstance(result_type, ir.IntType):
# 如果指针指向 i8加载字符值
if isinstance(OrelseVal.type.pointee, ir.IntType) and OrelseVal.type.pointee.width == 8:
OrelseVal = Gen._load(OrelseVal, name="ifexp_orelse_Load_char")
else:
OrelseVal = Gen.builder.ptrtoint(OrelseVal, result_type, name="ifexp_orelse_ptr2int")
Gen.builder.branch(MergeBB)
Gen.builder.position_at_start(MergeBB)
result_phi: Any = Gen.builder.phi(result_type, name="ifexp.result")
result_phi.add_incoming(BodyVal, ThenEndBB)
result_phi.add_incoming(OrelseVal, ElseEndBB)
return result_phi
def _HandleNamedExprLlvm(self, Node: ast.NamedExpr) -> ir.Value:
Gen: LlvmCodeGenerator = self.Trans.LlvmGen
ValueVal: Any = self.HandleExprLlvm(Node.value)
if not ValueVal:
return None
TargetName: str = Node.target.id
if TargetName in Gen._reg_values:
del Gen._reg_values[TargetName]
if TargetName in Gen.variables and Gen.variables[TargetName] is not None:
VarPtr: Any = Gen.variables[TargetName]
if VarPtr.type.pointee == ValueVal.type:
Gen._store(ValueVal, VarPtr)
else:
NewVar: Any = Gen._alloca(ValueVal.type, name=TargetName)
Gen._store(ValueVal, NewVar)
Gen.variables[TargetName] = NewVar
else:
NewVar = Gen._alloca(ValueVal.type, name=TargetName)
Gen._store(ValueVal, NewVar)
Gen.variables[TargetName] = NewVar
Gen._reg_values[TargetName] = ValueVal
return ValueVal