249 lines
14 KiB
Python
249 lines
14 KiB
Python
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 |