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