Files
TransPyC/lib/core/Handles/HandlesExprLambda.py
2026-07-18 19:25:40 +08:00

249 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
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