from __future__ import annotations from typing import TYPE_CHECKING if TYPE_CHECKING: from lib.core.translator import Translator from lib.core.Handles.HandlesBase import BaseHandle import ast import llvmlite.ir as ir class ForHandle(BaseHandle): def _HandleForLlvm(self, Node): Gen = self.Trans.LlvmGen Gen._set_node_info(Node, "For") if not isinstance(Node.target, ast.Name): return TargetName = Node.target.id StartVal = None EndVal = None StepVal = None IsRange = False CondOp = '<' if isinstance(Node.iter, ast.Call): CallFunc = Node.iter.func if isinstance(CallFunc, ast.Name) and CallFunc.id == 'range': IsRange = True RangeArgs = Node.iter.args RangeEndArg = RangeArgs[1] if len(RangeArgs) >= 2 else RangeArgs[0] if len(RangeArgs) >= 1 else None OptimizedEnd = False if len(RangeArgs) >= 2 and isinstance(RangeEndArg, ast.BinOp): if isinstance(RangeEndArg.op, ast.Add): if (isinstance(RangeEndArg.right, ast.Constant) and RangeEndArg.right.value == 1): EndVal = self.HandleExprLlvm(RangeEndArg.left) CondOp = '<=' OptimizedEnd = True elif (isinstance(RangeEndArg.left, ast.Constant) and RangeEndArg.left.value == 1): EndVal = self.HandleExprLlvm(RangeEndArg.right) CondOp = '<=' OptimizedEnd = True elif isinstance(RangeEndArg.op, ast.Sub): if isinstance(RangeEndArg.right, ast.Constant) and RangeEndArg.right.value == 1: EndVal = self.HandleExprLlvm(RangeEndArg.left) CondOp = '<' OptimizedEnd = True if not OptimizedEnd: if len(RangeArgs) == 1: EndVal = self.HandleExprLlvm(RangeArgs[0]) else: EndVal = self.HandleExprLlvm(RangeEndArg) if len(RangeArgs) == 1: StartVal = ir.Constant(ir.IntType(32), 0) elif len(RangeArgs) >= 2: StartVal = self.HandleExprLlvm(RangeArgs[0]) if len(RangeArgs) >= 3: StepVal = self.HandleExprLlvm(RangeArgs[2]) # Check if step is a negative constant — need to flip condition step_arg = RangeArgs[2] if isinstance(step_arg, ast.UnaryOp) and isinstance(step_arg.op, ast.USub): if isinstance(step_arg.operand, ast.Constant) and isinstance(step_arg.operand.value, (int, float)): if step_arg.operand.value > 0: CondOp = '>' elif isinstance(step_arg, ast.Constant) and isinstance(step_arg.value, (int, float)): if step_arg.value < 0: CondOp = '>' if not EndVal: return if not isinstance(EndVal.type, ir.IntType): if isinstance(EndVal.type, ir.PointerType) and isinstance(EndVal.type.pointee, ir.IntType): EndVal = Gen._load(EndVal, name="load_end") else: EndVal = Gen.builder.ptrtoint(EndVal, ir.IntType(64), name="end") if StartVal and not isinstance(StartVal.type, ir.IntType): if isinstance(StartVal.type, ir.PointerType) and isinstance(StartVal.type.pointee, ir.IntType): StartVal = Gen._load(StartVal, name="load_start") else: StartVal = Gen.builder.ptrtoint(StartVal, ir.IntType(64), name="start") else: StartVal = ir.Constant(ir.IntType(64), 0) if not StartVal: StartVal = ir.Constant(EndVal.type, 0) if StepVal is None: StepVal = ir.Constant(EndVal.type, 1) if not IsRange: IterClassName = self.Trans.ExprHandler._get_var_class(Node.iter, Gen) if not IterClassName and isinstance(Node.iter, ast.Call) and isinstance(Node.iter.func, ast.Name): IterClassName = Node.iter.func.id if IterClassName not in Gen.structs: IterClassName = None if IterClassName and Gen._has_function(f'{IterClassName}.__iter__') and Gen._has_function(f'{IterClassName}.__next__'): self._HandleForIterLlvm(Node, IterClassName) return IterVal = self.HandleExprLlvm(Node.iter) if not IterVal: return if isinstance(IterVal.type, ir.PointerType) and isinstance(IterVal.type.pointee, ir.IntType) and IterVal.type.pointee.width == 8: self._HandleForStrLlvm(Node, IterVal) return if isinstance(IterVal.type, ir.PointerType) and isinstance(IterVal.type.pointee, ir.ArrayType): self._HandleForListLlvm(Node, IterVal) return StartVal = ir.Constant(ir.IntType(32), 0) EndVal = IterVal if not isinstance(EndVal.type, ir.IntType): EndVal = Gen.builder.ptrtoint(EndVal, ir.IntType(64), name="end") StartVal = ir.Constant(ir.IntType(64), 0) StepVal = ir.Constant(EndVal.type, 1) CondOp = '<' if TargetName in Gen._reg_values: del Gen._reg_values[TargetName] if TargetName in Gen.variables and Gen.variables[TargetName] is not None: LoopVar = Gen.variables[TargetName] if isinstance(LoopVar.type, ir.PointerType) and isinstance(LoopVar.type.pointee, ir.IntType) and LoopVar.type.pointee == EndVal.type and isinstance(EndVal.type, ir.IntType) and StartVal.type == LoopVar.type.pointee and EndVal.type == StartVal.type: Gen._store(StartVal, LoopVar) else: NewVar = Gen._alloca_entry(EndVal.type, name=TargetName) if StartVal.type != EndVal.type: if isinstance(EndVal.type, ir.IntType) and isinstance(StartVal.type, ir.IntType): if StartVal.type.width < EndVal.type.width: StartVal = Gen.builder.zext(StartVal, EndVal.type, name="zext_loop_start") elif StartVal.type.width > EndVal.type.width: StartVal = Gen.builder.trunc(StartVal, EndVal.type, name="trunc_loop_start") Gen._store(StartVal, NewVar) Gen.variables[TargetName] = NewVar LoopVar = NewVar else: LoopVar = Gen._alloca_entry(EndVal.type, name=TargetName) Gen.variables[TargetName] = LoopVar Gen._store(StartVal, LoopVar) if StepVal and isinstance(StepVal.type, ir.IntType) and isinstance(EndVal.type, ir.IntType): if StepVal.type.width < EndVal.type.width: StepVal = Gen.builder.zext(StepVal, EndVal.type, name="zext_loop_step") elif StepVal.type.width > EndVal.type.width: StepVal = Gen.builder.trunc(StepVal, EndVal.type, name="trunc_loop_step") Gen._record_var_signedness(TargetName, 'int') CondBB = Gen.func.append_basic_block(name="for.cond") BodyBB = Gen.func.append_basic_block(name="for.body") StepBB = Gen.func.append_basic_block(name="for.step") ElseBB = Gen.func.append_basic_block(name="for.else") if Node.orelse else None AfterBB = Gen.func.append_basic_block(name="for.end") Gen.loop_break_targets.append(AfterBB) Gen.loop_continue_targets.append(StepBB) Gen.builder.branch(CondBB) Gen.builder.position_at_start(CondBB) LoopVal = Gen._load(LoopVar, name=TargetName) is_unsigned = Gen._is_var_unsigned(TargetName) if is_unsigned: Cond = Gen.builder.icmp_unsigned(CondOp, LoopVal, EndVal, name="forcond") else: Cond = Gen.builder.icmp_signed(CondOp, LoopVal, EndVal, name="forcond") Gen.builder.cbranch(Cond, BodyBB, ElseBB if ElseBB else AfterBB) Gen.builder.position_at_start(BodyBB) CachedLoopVal = Gen._load(LoopVar, name=TargetName) Gen._reg_values[TargetName] = CachedLoopVal self.HandleBodyLlvm(Node.body) StepLoopVal = CachedLoopVal if TargetName in Gen._reg_values and Gen._reg_values[TargetName] is CachedLoopVal: del Gen._reg_values[TargetName] else: StepLoopVal = Gen._load(LoopVar, name=TargetName) if not Gen.builder.block.is_terminated: Gen.builder.branch(StepBB) Gen.builder.position_at_start(StepBB) StepLoopVal2 = Gen._load(LoopVar, name=TargetName) NextVal = Gen.builder.add(StepLoopVal2, StepVal, name="next") Gen._store(NextVal, LoopVar) Gen.builder.branch(CondBB) Gen.loop_break_targets.pop() Gen.loop_continue_targets.pop() if ElseBB: Gen.builder.position_at_start(ElseBB) self.HandleBodyLlvm(Node.orelse) if not Gen.builder.block.is_terminated: Gen.builder.branch(AfterBB) Gen.builder.position_at_start(AfterBB) def _HandleGlobalLlvm(self, Node): Gen = self.Trans.LlvmGen for name in Node.names: Gen.global_vars.add(name) if name in Gen.module.globals: Gen.variables[name] = Gen.module.globals[name] if name in Gen._reg_values: del Gen._reg_values[name] def _HandleNonlocalLlvm(self, Node): Gen = self.Trans.LlvmGen for name in Node.names: if name in Gen._reg_values and name not in Gen.variables: OldVal = Gen._reg_values[name] var = Gen._alloca_entry(OldVal.type, name=name) Gen._store(OldVal, var) Gen.variables[name] = var del Gen._reg_values[name] def _collect_implicit_nonlocal(self, Node, Gen): param_names = set() for arg in Node.args.args: param_names.add(arg.arg) local_names = set() for child in ast.walk(Node): if isinstance(child, ast.Assign): for target in child.targets: if isinstance(target, ast.Name): local_names.add(target.id) elif isinstance(target, ast.Tuple): for elt in target.elts: if isinstance(elt, ast.Name): local_names.add(elt.id) elif isinstance(child, ast.AugAssign): if isinstance(child.target, ast.Name): local_names.add(child.target.id) elif isinstance(child, ast.AnnAssign): if isinstance(child.target, ast.Name): local_names.add(child.target.id) elif isinstance(child, ast.For): if isinstance(child.target, ast.Name): local_names.add(child.target.id) read_names = set() for child in ast.walk(Node): if isinstance(child, ast.Name) and isinstance(child.ctx, ast.Load): name = child.id if name not in param_names and name not in local_names: read_names.add(name) implicit = [] for name in read_names: if name in Gen.variables and Gen.variables[name] is not None: implicit.append(name) elif name in Gen._reg_values: implicit.append(name) for called_name in read_names: if called_name in Gen.nonlocal_params: for nv, _ in Gen.nonlocal_params[called_name]: if nv not in param_names and nv not in local_names: if nv not in read_names: if nv in Gen.variables and Gen.variables[nv] is not None: implicit.append(nv) elif nv in Gen._reg_values: implicit.append(nv) return implicit def _collect_all_nonlocal_names(self, Node): """Recursively collect all nonlocal variable names from a function and its nested functions.""" names = set() for stmt in Node.body: if isinstance(stmt, ast.Nonlocal): for name in stmt.names: names.add(name) elif isinstance(stmt, ast.FunctionDef): # Recursively collect from nested functions sub_names = self._collect_all_nonlocal_names(stmt) names.update(sub_names) return names def _HandleNestedFunctionLlvm(self, Node): Gen = self.Trans.LlvmGen nonlocal_var_names = [] for stmt in Node.body: if isinstance(stmt, ast.Nonlocal): for name in stmt.names: nonlocal_var_names.append(name) # Pre-scan ALL nested functions recursively for their nonlocal needs (passthrough) for stmt in Node.body: if isinstance(stmt, ast.FunctionDef): sub_nonlocal_names = self._collect_all_nonlocal_names(stmt) for name in sub_nonlocal_names: if name not in nonlocal_var_names: nonlocal_var_names.append(name) implicit_names = self._collect_implicit_nonlocal(Node, Gen) for name in implicit_names: if name not in nonlocal_var_names: nonlocal_var_names.append(name) extra_params = [] for var_name in nonlocal_var_names: if var_name in Gen.variables and Gen.variables[var_name] is not None: var = Gen.variables[var_name] extra_params.append((var_name, var.type)) elif var_name in Gen._reg_values: OldVal = Gen._reg_values[var_name] var = Gen._alloca_entry(OldVal.type, name=var_name) Gen._store(OldVal, var) Gen.variables[var_name] = var del Gen._reg_values[var_name] extra_params.append((var_name, var.type)) else: # Search outer scope stack for the variable found_var = None for scope in reversed(Gen._variable_scope_stack): if var_name in scope and scope[var_name] is not None: found_var = scope[var_name] break if found_var is not None: # Add to current scope so it can be passed to nested functions Gen.variables[var_name] = found_var extra_params.append((var_name, found_var.type)) else: var = Gen._alloca_entry(ir.IntType(32), name=var_name) Gen.variables[var_name] = var extra_params.append((var_name, var.type)) if extra_params: Gen.nonlocal_params[Node.name] = [(name, ptr_type) for name, ptr_type in extra_params] # Push current variables to scope stack before compiling nested function Gen._variable_scope_stack.append(Gen.variables.copy()) saved_builder = Gen.builder saved_func = Gen.func saved_variables = Gen.variables.copy() saved_reg_values = Gen._reg_values.copy() saved_global_vars = Gen.global_vars.copy() saved_var_signedness = Gen.var_signedness.copy() saved_var_type_info = Gen.var_type_info.copy() saved_var_type_assignments = Gen.var_type_assignments.copy() saved_CurrentCReturnTypes = getattr(self.Trans, 'CurrentCReturnTypes', None) saved_CurrentCpythonObjectClass = getattr(self.Trans, '_CurrentCpythonObjectClass', None) saved_VarScopes = len(self.Trans.VarScopes) self.Trans.FunctionHandler._EmitFunctionLlvm(Node, Gen, extra_params=extra_params if extra_params else None) Gen.builder = saved_builder Gen.func = saved_func Gen.variables = saved_variables Gen._reg_values = saved_reg_values Gen.global_vars = saved_global_vars Gen.var_signedness = saved_var_signedness Gen.var_type_info = saved_var_type_info Gen.var_type_assignments = saved_var_type_assignments self.Trans.CurrentCReturnTypes = saved_CurrentCReturnTypes self.Trans._CurrentCpythonObjectClass = saved_CurrentCpythonObjectClass while len(self.Trans.VarScopes) > saved_VarScopes: self.Trans.VarScopes.pop() # Pop scope stack if Gen._variable_scope_stack: Gen._variable_scope_stack.pop() def _HandleForListLlvm(self, Node, ArrPtr): Gen = self.Trans.LlvmGen if not isinstance(Node.target, ast.Name): return TargetName = Node.target.id ElemType = ArrPtr.type.pointee.element ArrayCount = ArrPtr.type.pointee.count IdxVal = Gen._alloca_entry(ir.IntType(32), name=f"arr_idx") Gen.builder.store(ir.Constant(ir.IntType(32), 0), IdxVal) if TargetName in Gen._reg_values: del Gen._reg_values[TargetName] LoopVar = Gen._alloca_entry(ElemType, name=TargetName) Gen.variables[TargetName] = LoopVar Gen._store(ir.Constant(ElemType, None), LoopVar) CondBB = Gen.func.append_basic_block(name="arrfor.cond") BodyBB = Gen.func.append_basic_block(name="arrfor.body") StepBB = Gen.func.append_basic_block(name="arrfor.step") ElseBB = Gen.func.append_basic_block(name="arrfor.else") if Node.orelse else None AfterBB = Gen.func.append_basic_block(name="arrfor.end") Gen.loop_break_targets.append(AfterBB) Gen.loop_continue_targets.append(StepBB) Gen.builder.branch(CondBB) Gen.builder.position_at_start(CondBB) CurIdx = Gen._load(IdxVal, name="cur_idx") IsDone = Gen.builder.icmp_signed('<', CurIdx, ir.Constant(ir.IntType(32), ArrayCount), name="arr_in_bounds") Gen.builder.cbranch(IsDone, BodyBB, ElseBB if ElseBB else AfterBB) Gen.builder.position_at_start(BodyBB) ElemPtr = Gen.builder.gep(ArrPtr, [ir.Constant(ir.IntType(32), 0), CurIdx], name="arr_elem_ptr") ElemVal = Gen._load(ElemPtr, name="arr_elem_val") Gen._store(ElemVal, LoopVar) Gen._record_var_signedness(TargetName, 'ptr') Gen._reg_values[TargetName] = ElemVal self.HandleBodyLlvm(Node.body) if TargetName in Gen._reg_values and Gen._reg_values[TargetName] is ElemVal: del Gen._reg_values[TargetName] if not Gen.builder.block.is_terminated: Gen.builder.branch(StepBB) Gen.builder.position_at_start(StepBB) CurIdx2 = Gen._load(IdxVal, name="cur_idx2") NextIdx = Gen.builder.add(CurIdx2, ir.Constant(ir.IntType(32), 1), name="next_idx") Gen._store(NextIdx, IdxVal) Gen.builder.branch(CondBB) Gen.loop_break_targets.pop() Gen.loop_continue_targets.pop() if ElseBB: Gen.builder.position_at_start(ElseBB) self.HandleBodyLlvm(Node.orelse) if not Gen.builder.block.is_terminated: Gen.builder.branch(AfterBB) Gen.builder.position_at_start(AfterBB) def _HandleForStrLlvm(self, Node, StrVal): Gen = self.Trans.LlvmGen if not isinstance(Node.target, ast.Name): return TargetName = Node.target.id Gen._unregister_temp_ptr(StrVal) PtrVal = Gen.builder.alloca(ir.IntType(8).as_pointer(), name=f"str_ptr_copy") Gen.builder.store(StrVal, PtrVal) CharType = ir.IntType(8) if TargetName in Gen._reg_values: del Gen._reg_values[TargetName] LoopVar = Gen._alloca_entry(CharType, name=TargetName) Gen.variables[TargetName] = LoopVar Gen._store(ir.Constant(CharType, 0), LoopVar) CondBB = Gen.func.append_basic_block(name="strfor.cond") BodyBB = Gen.func.append_basic_block(name="strfor.body") StepBB = Gen.func.append_basic_block(name="strfor.step") ElseBB = Gen.func.append_basic_block(name="strfor.else") if Node.orelse else None AfterBB = Gen.func.append_basic_block(name="strfor.end") Gen.loop_break_targets.append(AfterBB) Gen.loop_continue_targets.append(StepBB) Gen.builder.branch(CondBB) Gen.builder.position_at_start(CondBB) CurrentPtr = Gen._load(PtrVal, name="current_ptr") RawChar = Gen._load(CurrentPtr, name="raw_char") NullCond = Gen.builder.icmp_signed('==', RawChar, ir.Constant(CharType, 0), name="is_null") Gen.builder.cbranch(NullCond, ElseBB if ElseBB else AfterBB, BodyBB) Gen.builder.position_at_start(BodyBB) Gen._store(RawChar, LoopVar) Gen._record_var_signedness(TargetName, 'char') Gen._reg_values[TargetName] = RawChar self.HandleBodyLlvm(Node.body) if TargetName in Gen._reg_values and Gen._reg_values[TargetName] is RawChar: del Gen._reg_values[TargetName] if not Gen.builder.block.is_terminated: Gen.builder.branch(StepBB) Gen.builder.position_at_start(StepBB) CurrentPtr2 = Gen._load(PtrVal, name="current_ptr2") NextPtr = Gen.builder.gep(CurrentPtr2, [ir.Constant(ir.IntType(32), 1)], name="next_ptr") Gen._store(NextPtr, PtrVal) Gen.builder.branch(CondBB) Gen.loop_break_targets.pop() Gen.loop_continue_targets.pop() if ElseBB: Gen.builder.position_at_start(ElseBB) self.HandleBodyLlvm(Node.orelse) if not Gen.builder.block.is_terminated: Gen.builder.branch(AfterBB) Gen.builder.position_at_start(AfterBB) def _HandleForIterLlvm(self, Node, ClassName): Gen = self.Trans.LlvmGen if not isinstance(Node.target, ast.Name): return TargetName = Node.target.id IterVal = self.HandleExprLlvm(Node.iter) if not IterVal: return Gen._unregister_temp_ptr(IterVal) if isinstance(IterVal.type, ir.PointerType) and isinstance(IterVal.type.pointee, ir.IntType) and IterVal.type.pointee.width == 8: IterVal = Gen.builder.bitcast(IterVal, ir.PointerType(Gen.structs[ClassName]), name=f"cast_{ClassName}") IterCall = Gen._get_function(f'{ClassName}.__iter__') IterResult = Gen.builder.call(IterCall, [IterVal], name=f"call_{ClassName}.__iter__") Gen._unregister_temp_ptr(IterResult) if isinstance(IterResult.type, ir.PointerType) and isinstance(IterResult.type.pointee, ir.IntType) and IterResult.type.pointee.width == 8: IterResult = Gen.builder.bitcast(IterResult, ir.PointerType(Gen.structs[ClassName]), name=f"cast_iter_{ClassName}") NextCall = Gen._get_function(f'{ClassName}.__next__') StopFlagPtr = Gen._alloca_entry(ir.IntType(1), name="stop_iter_flag") Gen.builder.store(ir.Constant(ir.IntType(1), 0), StopFlagPtr) NextReturnType = NextCall.function_type.return_type if TargetName in Gen._reg_values: del Gen._reg_values[TargetName] LoopVar = Gen._alloca_entry(NextReturnType, name=TargetName) Gen.variables[TargetName] = LoopVar CondBB = Gen.func.append_basic_block(name="iter.cond") BodyBB = Gen.func.append_basic_block(name="iter.body") StepBB = Gen.func.append_basic_block(name="iter.step") ElseBB = Gen.func.append_basic_block(name="iter.else") if Node.orelse else None AfterBB = Gen.func.append_basic_block(name="iter.end") Gen.loop_break_targets.append(AfterBB) Gen.loop_continue_targets.append(StepBB) Gen.builder.branch(CondBB) Gen.builder.position_at_start(CondBB) Gen.builder.store(ir.Constant(ir.IntType(1), 0), StopFlagPtr) NextCallArgs = [IterResult, StopFlagPtr] if NextCall and len(NextCall.function_type.args) > len(NextCallArgs): last_param_type = NextCall.function_type.args[-1] if (isinstance(last_param_type, ir.PointerType) and isinstance(last_param_type.pointee, ir.PointerType) and isinstance(last_param_type.pointee.pointee, ir.IntType) and last_param_type.pointee.pointee.width == 8): null_msg = Gen._alloca_entry(ir.PointerType(ir.IntType(8)), name="eh_msg_out_null") NextCallArgs.append(null_msg) NextVal = Gen.builder.call(NextCall, NextCallArgs, name=f"call_{ClassName}.__next__") StopFlag = Gen.builder.load(StopFlagPtr, name="stop_flag") Cond = Gen.builder.icmp_signed('==', StopFlag, ir.Constant(ir.IntType(1), 0), name="iter_cond") Gen.builder.cbranch(Cond, BodyBB, ElseBB if ElseBB else AfterBB) Gen.builder.position_at_start(BodyBB) Gen._store(NextVal, LoopVar) Gen._reg_values[TargetName] = NextVal self.HandleBodyLlvm(Node.body) if TargetName in Gen._reg_values and Gen._reg_values[TargetName] is NextVal: del Gen._reg_values[TargetName] if not Gen.builder.block.is_terminated: Gen.builder.branch(StepBB) Gen.builder.position_at_start(StepBB) Gen.builder.branch(CondBB) Gen.loop_break_targets.pop() Gen.loop_continue_targets.pop() if ElseBB: Gen.builder.position_at_start(ElseBB) self.HandleBodyLlvm(Node.orelse) if not Gen.builder.block.is_terminated: Gen.builder.branch(AfterBB) Gen.builder.position_at_start(AfterBB)