Files
TransPyC/lib/core/Handles/HandlesFor.py

1073 lines
58 KiB
Python
Raw 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
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: ast.For) -> None:
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
Gen._set_node_info(Node, "For")
# zip() 内置for a, b in zip(iter1, iter2):
if (isinstance(Node.iter, ast.Call)
and isinstance(Node.iter.func, ast.Name)
and Node.iter.func.id == 'zip'
and isinstance(Node.target, ast.Tuple)):
self._HandleForZipLlvm(Node)
return
# enumerate() 内置for idx, val in enumerate(iterable):
if (isinstance(Node.iter, ast.Call)
and isinstance(Node.iter.func, ast.Name)
and Node.iter.func.id == 'enumerate'
and isinstance(Node.target, ast.Tuple)
and Node.iter.args):
self._HandleForEnumerateLlvm(Node)
return
if not isinstance(Node.target, ast.Name):
return
TargetName: str = Node.target.id
StartVal: ir.Value | None = None
EndVal: ir.Value | None = None
StepVal: ir.Value | None = None
IsRange: bool = False
CondOp: str = '<'
if isinstance(Node.iter, ast.Call):
CallFunc: ast.expr = Node.iter.func
if isinstance(CallFunc, ast.Name) and CallFunc.id == 'range':
IsRange = True
RangeArgs: list[ast.expr] = Node.iter.args
RangeEndArg: ast.expr | None = RangeArgs[1] if len(RangeArgs) >= 2 else RangeArgs[0] if len(RangeArgs) >= 1 else None
OptimizedEnd: bool = 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: ast.expr = 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:
# 检测 c.Addr(x) 模式按原始变量类型迭代for i64 in i64* 语义)
if (isinstance(Node.iter, ast.Call)
and isinstance(Node.iter.func, ast.Attribute)
and isinstance(Node.iter.func.value, ast.Name)
and Node.iter.func.value.id == 'c'
and Node.iter.func.attr == 'Addr'
and Node.iter.args):
AddrArg: ast.expr = Node.iter.args[0]
if isinstance(AddrArg, ast.Name) and AddrArg.id in Gen.variables:
_var_ptr: ir.Value | None = Gen.variables[AddrArg.id]
if _var_ptr is not None and isinstance(_var_ptr.type, ir.PointerType):
_pointee: ir.Type = _var_ptr.type.pointee
# 标量整数变量
if isinstance(_pointee, ir.IntType):
if _pointee.width == 8:
self._HandleForStrLlvm(Node, _var_ptr)
return
self._HandleForPtrLlvm(Node, _var_ptr, _pointee)
return
# 数组变量:取元素类型
if isinstance(_pointee, ir.ArrayType) and isinstance(_pointee.element, ir.IntType):
_elem_type: ir.IntType = _pointee.element
if _elem_type.width == 8:
self._HandleForStrLlvm(Node, _var_ptr)
return
_typed_ptr: ir.Value = Gen.builder.bitcast(_var_ptr, ir.PointerType(_elem_type), name="addr_arr_elem_cast")
self._HandleForPtrLlvm(Node, _typed_ptr, _elem_type)
return
IterClassName: str | None = 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: ir.Value | None = self.HandleExprLlvm(Node.iter)
if not IterVal:
return
# 回退:通过 LLVM 类型查找结构体类名,检测 __iter__/__next__
if not IterClassName and isinstance(IterVal.type, ir.PointerType):
_pointee: ir.Type = IterVal.type.pointee
if isinstance(_pointee, (ir.IdentifiedStructType, ir.LiteralStructType)):
_found: tuple[str, ir.Type] | None = Gen.find_struct_by_pointee(_pointee)
if _found:
_cn: str = _found[0]
if not (Gen._has_function(f'{_cn}.__iter__') and Gen._has_function(f'{_cn}.__next__')):
# temp stub 可能不包含泛型特化方法,尝试从 output stub 按需加载
self.Trans.ImportHandler._TryLoadFuncDeclsFromOutputStub(_cn, Gen)
if Gen._has_function(f'{_cn}.__iter__') and Gen._has_function(f'{_cn}.__next__'):
Gen._UnregisterTempPtr(IterVal)
self._HandleForIterLlvm(Node, _cn)
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]
LoopVar: ir.AllocaInstr
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: ir.AllocaInstr = Gen._allocaEntry(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._allocaEntry(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: ir.Block = Gen.func.append_basic_block(name="for.cond")
BodyBB: ir.Block = Gen.func.append_basic_block(name="for.body")
StepBB: ir.Block = Gen.func.append_basic_block(name="for.step")
ElseBB: ir.Block | None = Gen.func.append_basic_block(name="for.else") if Node.orelse else None
AfterBB: ir.Block = 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: ir.Value = Gen._load(LoopVar, name=TargetName)
is_unsigned: bool = Gen._is_var_unsigned(TargetName)
Cond: ir.Value
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: ir.Value = Gen._load(LoopVar, name=TargetName)
Gen._reg_values[TargetName] = CachedLoopVal
self.HandleBodyLlvm(Node.body)
StepLoopVal: ir.Value = 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: ir.Value = Gen._load(LoopVar, name=TargetName)
NextVal: ir.Value = 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: ast.Global) -> None:
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
self._RegisterGlobalNames(Node.names, Gen, clear_reg_values=True)
def _RegisterGlobalNames(self, names: list[str], Gen: "Translator.LlvmGen", clear_reg_values: bool = False) -> None:
for name in names:
Gen.global_vars.add(name)
if name in Gen.module.globals:
Gen.variables[name] = Gen.module.globals[name]
if clear_reg_values and name in Gen._reg_values:
del Gen._reg_values[name]
def _HandleNonlocalLlvm(self, Node: ast.Nonlocal) -> None:
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
for name in Node.names:
if name in Gen._reg_values and name not in Gen.variables:
OldVal: ir.Value = Gen._reg_values[name]
var: ir.AllocaInstr = Gen._allocaEntry(OldVal.type, name=name)
Gen._store(OldVal, var)
Gen.variables[name] = var
del Gen._reg_values[name]
def _collect_implicit_nonlocal(self, Node: ast.FunctionDef, Gen: "Translator.LlvmGen") -> list[str]:
param_names: set[str] = set()
for arg in Node.args.args:
param_names.add(arg.arg)
local_names: set[str] = 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[str] = set()
for child in ast.walk(Node):
if isinstance(child, ast.Name) and isinstance(child.ctx, ast.Load):
name: str = child.id
if name not in param_names and name not in local_names:
read_names.add(name)
implicit: list[str] = []
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: ast.FunctionDef) -> set[str]:
"""Recursively collect all nonlocal variable names from a function and its nested functions."""
names: set[str] = 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: set[str] = self._collect_all_nonlocal_names(stmt)
names.update(sub_names)
return names
def _HandleNestedFunctionLlvm(self, Node: ast.FunctionDef) -> None:
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
nonlocal_var_names: list[str] = []
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: set[str] = 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: list[str] = self._collect_implicit_nonlocal(Node, Gen)
for name in implicit_names:
if name not in nonlocal_var_names:
nonlocal_var_names.append(name)
# 堆分配 nonlocal 变量:避免外层函数返回后栈帧释放导致悬垂指针
# 闭包返回模式return nested_fn要求捕获变量在堆上存活
malloc_fn_type: ir.FunctionType = ir.FunctionType(ir.IntType(8).as_pointer(), [ir.IntType(64)])
malloc_func: ir.Function = Gen._get_or_declare_func('malloc', malloc_fn_type)
for var_name in nonlocal_var_names:
if var_name in Gen.variables and Gen.variables[var_name] is not None:
old_var: ir.Value = Gen.variables[var_name]
if isinstance(old_var.type, ir.PointerType):
elem_type: ir.Type = old_var.type.pointee
# 计算元素大小(字节)
elem_size: int = 8 # 默认指针大小
if isinstance(elem_type, ir.IntType):
elem_size = max(1, elem_type.width // 8)
elif isinstance(elem_type, ir.FloatType):
elem_size = 4
elif isinstance(elem_type, ir.DoubleType):
elem_size = 8
elif isinstance(elem_type, ir.PointerType):
elem_size = 8
# malloc 堆存储
heap_raw: ir.Value = Gen.builder.call(malloc_func, [ir.Constant(ir.IntType(64), elem_size)], name=f"heap_{var_name}")
heap_ptr: ir.Value = Gen.builder.bitcast(heap_raw, ir.PointerType(elem_type), name=f"heap_{var_name}_ptr")
# 复制栈值到堆
old_val: ir.Value = Gen._load(old_var, name=f"load_old_{var_name}")
Gen._store(old_val, heap_ptr)
# 更新 Gen.variables 使外层和嵌套函数都使用堆指针
Gen.variables[var_name] = heap_ptr
extra_params: list[tuple[str, ir.Type]] = []
for var_name in nonlocal_var_names:
if var_name in Gen.variables and Gen.variables[var_name] is not None:
var: ir.Value = Gen.variables[var_name]
extra_params.append((var_name, var.type))
elif var_name in Gen._reg_values:
OldVal: ir.Value = Gen._reg_values[var_name]
var: ir.AllocaInstr = Gen._allocaEntry(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: ir.Value | None = 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._allocaEntry(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: ir.IRBuilder = Gen.builder
saved_func: ir.Function = Gen.func
saved_variables: dict[str, ir.Value] = Gen.variables.copy()
saved_reg_values: dict[str, ir.Value] = Gen._reg_values.copy()
saved_global_vars: set[str] = Gen.global_vars.copy()
saved_var_signedness: dict[str, str] = Gen.var_signedness.copy()
saved_var_type_info: dict[str, dict[str, str]] = Gen.var_type_info.copy()
saved_var_type_assignments: dict[str, list[dict[str, str]]] = Gen.var_type_assignments.copy()
saved_CurrentCReturnTypes: object | None = getattr(self.Trans, 'CurrentCReturnTypes', None)
saved_CurrentCpythonObjectClass: str | None = getattr(self.Trans, '_CurrentCpythonObjectClass', None)
saved_VarScopes: int = 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: ast.For, ArrPtr: ir.Value) -> None:
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
if not isinstance(Node.target, ast.Name):
return
TargetName: str = Node.target.id
ElemType: ir.Type = ArrPtr.type.pointee.element
ArrayCount: int = ArrPtr.type.pointee.count
IdxVal: ir.AllocaInstr = Gen._allocaEntry(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: ir.AllocaInstr = Gen._allocaEntry(ElemType, name=TargetName)
Gen.variables[TargetName] = LoopVar
Gen._store(ir.Constant(ElemType, None), LoopVar)
CondBB: ir.Block = Gen.func.append_basic_block(name="arrfor.cond")
BodyBB: ir.Block = Gen.func.append_basic_block(name="arrfor.body")
StepBB: ir.Block = Gen.func.append_basic_block(name="arrfor.step")
ElseBB: ir.Block | None = Gen.func.append_basic_block(name="arrfor.else") if Node.orelse else None
AfterBB: ir.Block = 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: ir.Value = Gen._load(IdxVal, name="cur_idx")
IsDone: ir.Value = 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: ir.Value = Gen.builder.gep(ArrPtr, [ir.Constant(ir.IntType(32), 0), CurIdx], name="arr_elem_ptr")
ElemVal: ir.Value = 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: ir.Value = Gen._load(IdxVal, name="cur_idx2")
NextIdx: ir.Value = 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: ast.For, StrVal: ir.Value) -> None:
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
if not isinstance(Node.target, ast.Name):
return
TargetName: str = Node.target.id
Gen._UnregisterTempPtr(StrVal)
PtrVal: ir.AllocaInstr = Gen.builder.alloca(ir.IntType(8).as_pointer(), name=f"str_ptr_copy")
Gen.builder.store(StrVal, PtrVal)
CharType: ir.IntType = ir.IntType(8)
if TargetName in Gen._reg_values:
del Gen._reg_values[TargetName]
LoopVar: ir.AllocaInstr = Gen._allocaEntry(CharType, name=TargetName)
Gen.variables[TargetName] = LoopVar
Gen._store(ir.Constant(CharType, 0), LoopVar)
CondBB: ir.Block = Gen.func.append_basic_block(name="strfor.cond")
BodyBB: ir.Block = Gen.func.append_basic_block(name="strfor.body")
StepBB: ir.Block = Gen.func.append_basic_block(name="strfor.step")
ElseBB: ir.Block | None = Gen.func.append_basic_block(name="strfor.else") if Node.orelse else None
AfterBB: ir.Block = 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: ir.Value = Gen._load(PtrVal, name="current_ptr")
RawChar: ir.Value = Gen._load(CurrentPtr, name="raw_char")
NullCond: ir.Value = 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: ir.Value = Gen._load(PtrVal, name="current_ptr2")
NextPtr: ir.Value = 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 _HandleForPtrLlvm(self, Node: ast.For, PtrVal: ir.Value, ElemType: ir.IntType) -> None:
"""通用指针迭代:按 ElemType 大小逐元素迭代,终止条件为元素值 == 0
适用于 for i in c.Addr(u) where u: i32/i64 等非 i8 整数类型。
"""
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
if not isinstance(Node.target, ast.Name):
return
TargetName: str = Node.target.id
# 确保 PtrVal 是 ElemType*
TypedPtr: ir.Value = PtrVal
if isinstance(PtrVal.type, ir.PointerType) and PtrVal.type.pointee is not ElemType:
TypedPtr = Gen.builder.bitcast(PtrVal, ir.PointerType(ElemType), name="ptr_elem_cast")
Gen._UnregisterTempPtr(TypedPtr)
# 分配指针副本
PtrCopy: ir.AllocaInstr = Gen.builder.alloca(ir.PointerType(ElemType), name="ptr_copy")
Gen.builder.store(TypedPtr, PtrCopy)
if TargetName in Gen._reg_values:
del Gen._reg_values[TargetName]
LoopVar: ir.AllocaInstr = Gen._allocaEntry(ElemType, name=TargetName)
Gen.variables[TargetName] = LoopVar
Gen._store(ir.Constant(ElemType, 0), LoopVar)
CondBB: ir.Block = Gen.func.append_basic_block(name="ptrfor.cond")
BodyBB: ir.Block = Gen.func.append_basic_block(name="ptrfor.body")
StepBB: ir.Block = Gen.func.append_basic_block(name="ptrfor.step")
ElseBB: ir.Block | None = Gen.func.append_basic_block(name="ptrfor.else") if Node.orelse else None
AfterBB: ir.Block = Gen.func.append_basic_block(name="ptrfor.end")
Gen.loop_break_targets.append(AfterBB)
Gen.loop_continue_targets.append(StepBB)
Gen.builder.branch(CondBB)
Gen.builder.position_at_start(CondBB)
CurrentPtr: ir.Value = Gen._load(PtrCopy, name="current_ptr")
RawElem: ir.Value = Gen._load(CurrentPtr, name="raw_elem")
NullCond: ir.Value = Gen.builder.icmp_signed('==', RawElem, ir.Constant(ElemType, 0), name="is_null")
Gen.builder.cbranch(NullCond, ElseBB if ElseBB else AfterBB, BodyBB)
Gen.builder.position_at_start(BodyBB)
Gen._store(RawElem, LoopVar)
Gen._record_var_signedness(TargetName, 'int')
Gen._reg_values[TargetName] = RawElem
self.HandleBodyLlvm(Node.body)
if TargetName in Gen._reg_values and Gen._reg_values[TargetName] is RawElem:
del Gen._reg_values[TargetName]
if not Gen.builder.block.is_terminated:
Gen.builder.branch(StepBB)
Gen.builder.position_at_start(StepBB)
CurrentPtr2: ir.Value = Gen._load(PtrCopy, name="current_ptr2")
NextPtr: ir.Value = Gen.builder.gep(CurrentPtr2, [ir.Constant(ir.IntType(32), 1)], name="next_ptr")
Gen._store(NextPtr, PtrCopy)
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: ast.For, ClassName: str) -> None:
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
if not isinstance(Node.target, ast.Name):
return
TargetName: str = Node.target.id
IterVal: ir.Value | None = self.HandleExprLlvm(Node.iter)
if not IterVal:
return
Gen._UnregisterTempPtr(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: ir.Function = Gen._get_function(f'{ClassName}.__iter__')
IterResult: ir.Value = Gen.builder.call(IterCall, [IterVal], name=f"call_{ClassName}.__iter__")
Gen._UnregisterTempPtr(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: ir.Function = Gen._get_function(f'{ClassName}.__next__')
StopFlagPtr: ir.AllocaInstr = Gen._allocaEntry(ir.IntType(1), name="stop_iter_flag")
Gen.builder.store(ir.Constant(ir.IntType(1), 0), StopFlagPtr)
NextReturnType: ir.Type = NextCall.function_type.return_type
if TargetName in Gen._reg_values:
del Gen._reg_values[TargetName]
LoopVar: ir.AllocaInstr = Gen._allocaEntry(NextReturnType, name=TargetName)
Gen.variables[TargetName] = LoopVar
CondBB: ir.Block = Gen.func.append_basic_block(name="iter.cond")
BodyBB: ir.Block = Gen.func.append_basic_block(name="iter.body")
StepBB: ir.Block = Gen.func.append_basic_block(name="iter.step")
ElseBB: ir.Block | None = Gen.func.append_basic_block(name="iter.else") if Node.orelse else None
AfterBB: ir.Block = 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: list[ir.Value] = [IterResult, StopFlagPtr]
if NextCall and len(NextCall.function_type.args) > len(NextCallArgs):
last_param_type: ir.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: ir.AllocaInstr = Gen._allocaEntry(ir.PointerType(ir.IntType(8)), name="eh_msg_out_null")
NextCallArgs.append(null_msg)
NextVal: ir.Value = Gen.builder.call(NextCall, NextCallArgs, name=f"call_{ClassName}.__next__")
StopFlag: ir.Value = Gen.builder.load(StopFlagPtr, name="stop_flag")
Cond: ir.Value = 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)
# ===== zip() 内置支持 =====
def _HandleForZipLlvm(self, Node: ast.For) -> None:
"""for a, b in zip(iter1, iter2): — 多迭代器并行迭代
兼容 range、__iter__/__next__、字符串、数组、指针等所有迭代类型。
任一迭代器停止则退出循环。
"""
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
if not isinstance(Node.target, ast.Tuple):
return
Targets: list[ast.Name] = [t for t in Node.target.elts if isinstance(t, ast.Name)]
ZipArgs: list[ast.expr] = Node.iter.args
if len(Targets) != len(ZipArgs) or len(Targets) < 2:
return
# 为每个可迭代对象设置迭代状态
IterInfos: list[dict] = []
for i, arg in enumerate(ZipArgs):
target_name: str = Targets[i].id
info: dict | None = self._zip_setup_iter(arg, target_name, Gen)
if info is None:
return
IterInfos.append(info)
CondBB: ir.Block = Gen.func.append_basic_block(name="zip.cond")
BodyBB: ir.Block = Gen.func.append_basic_block(name="zip.body")
StepBB: ir.Block = Gen.func.append_basic_block(name="zip.step")
ElseBB: ir.Block | None = Gen.func.append_basic_block(name="zip.else") if Node.orelse else None
AfterBB: ir.Block = Gen.func.append_basic_block(name="zip.end")
Gen.loop_break_targets.append(AfterBB)
Gen.loop_continue_targets.append(StepBB)
Gen.builder.branch(CondBB)
# Cond 块:为每个迭代器获取下一个值和停止条件
Gen.builder.position_at_start(CondBB)
AnyStopped: ir.Value | None = None
for info in IterInfos:
val: ir.Value
stopped: ir.Value
val, stopped = self._zip_get_next(info, Gen)
info['last_value'] = val
if AnyStopped is None:
AnyStopped = stopped
else:
AnyStopped = Gen.builder.or_(AnyStopped, stopped, name="zip_any_stop")
Gen.builder.cbranch(AnyStopped, ElseBB if ElseBB else AfterBB, BodyBB)
# Body 块:赋值并执行循环体
Gen.builder.position_at_start(BodyBB)
for info in IterInfos:
Gen._store(info['last_value'], info['loop_var'])
Gen._reg_values[info['target_name']] = info['last_value']
self.HandleBodyLlvm(Node.body)
for info in IterInfos:
tn: str = info['target_name']
if tn in Gen._reg_values and Gen._reg_values[tn] is info['last_value']:
del Gen._reg_values[tn]
if not Gen.builder.block.is_terminated:
Gen.builder.branch(StepBB)
# Step 块:跳回 Cond推进已在 Cond 中完成)
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)
def _zip_setup_iter(self, arg: ast.expr, target_name: str, Gen: "Translator.LlvmGen") -> dict | None:
"""为 zip 参数设置迭代状态,返回包含迭代信息的 dict。"""
if target_name in Gen._reg_values:
del Gen._reg_values[target_name]
# Case 1: range(n) / range(start, end) / range(start, end, step)
if isinstance(arg, ast.Call) and isinstance(arg.func, ast.Name) and arg.func.id == 'range':
return self._zip_setup_range(arg, target_name, Gen)
# Case 2: c.Addr(x) — 指针/数组迭代
if (isinstance(arg, ast.Call)
and isinstance(arg.func, ast.Attribute)
and isinstance(arg.func.value, ast.Name)
and arg.func.value.id == 'c'
and arg.func.attr == 'Addr'
and arg.args):
AddrArg: ast.expr = arg.args[0]
if isinstance(AddrArg, ast.Name) and AddrArg.id in Gen.variables:
_var_ptr: ir.Value | None = Gen.variables[AddrArg.id]
if _var_ptr is not None and isinstance(_var_ptr.type, ir.PointerType):
_pointee: ir.Type = _var_ptr.type.pointee
if isinstance(_pointee, ir.IntType):
if _pointee.width == 8:
return self._zip_setup_str(_var_ptr, target_name, Gen)
return self._zip_setup_ptr(_var_ptr, target_name, _pointee, Gen)
if isinstance(_pointee, ir.ArrayType) and isinstance(_pointee.element, ir.IntType):
_elem_type: ir.IntType = _pointee.element
if _elem_type.width == 8:
return self._zip_setup_str(_var_ptr, target_name, Gen)
_typed_ptr: ir.Value = Gen.builder.bitcast(_var_ptr, ir.PointerType(_elem_type), name="zip_addr_cast")
return self._zip_setup_ptr(_typed_ptr, target_name, _elem_type, Gen)
# Case 3: __iter__/__next__ 类迭代
IterClassName: str | None = self.Trans.ExprHandler._get_var_class(arg, Gen)
if not IterClassName and isinstance(arg, ast.Call) and isinstance(arg.func, ast.Name):
IterClassName = arg.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__'):
return self._zip_setup_iter_class(arg, target_name, IterClassName, Gen)
# Case 4-6: 求值表达式后按类型分发
IterVal: ir.Value | None = self.HandleExprLlvm(arg)
if not IterVal:
return None
# 回退:通过 LLVM 类型查找结构体类名,检测 __iter__/__next__
if not IterClassName and isinstance(IterVal.type, ir.PointerType):
_pointee: ir.Type = IterVal.type.pointee
if isinstance(_pointee, (ir.IdentifiedStructType, ir.LiteralStructType)):
_found: tuple[str, ir.Type] | None = Gen.find_struct_by_pointee(_pointee)
if _found:
_cn: str = _found[0]
if not (Gen._has_function(f'{_cn}.__iter__') and Gen._has_function(f'{_cn}.__next__')):
self.Trans.ImportHandler._TryLoadFuncDeclsFromOutputStub(_cn, Gen)
if Gen._has_function(f'{_cn}.__iter__') and Gen._has_function(f'{_cn}.__next__'):
Gen._UnregisterTempPtr(IterVal)
return self._zip_setup_iter_class_from_val(IterVal, target_name, _cn, Gen)
if isinstance(IterVal.type, ir.PointerType) and isinstance(IterVal.type.pointee, ir.IntType) and IterVal.type.pointee.width == 8:
return self._zip_setup_str(IterVal, target_name, Gen)
if isinstance(IterVal.type, ir.PointerType) and isinstance(IterVal.type.pointee, ir.ArrayType):
return self._zip_setup_array(IterVal, target_name, Gen)
# Case 6: 整数计数器
return self._zip_setup_int(IterVal, target_name, Gen)
def _zip_setup_range(self, arg: ast.Call, target_name: str, Gen: "Translator.LlvmGen") -> dict | None:
RangeArgs: list[ast.expr] = arg.args
if len(RangeArgs) == 1:
StartVal: ir.Value = ir.Constant(ir.IntType(32), 0)
EndVal: ir.Value = self.HandleExprLlvm(RangeArgs[0])
StepVal: ir.Value = ir.Constant(ir.IntType(32), 1)
elif len(RangeArgs) == 2:
StartVal = self.HandleExprLlvm(RangeArgs[0])
EndVal = self.HandleExprLlvm(RangeArgs[1])
StepVal = ir.Constant(ir.IntType(32), 1)
elif len(RangeArgs) >= 3:
StartVal = self.HandleExprLlvm(RangeArgs[0])
EndVal = self.HandleExprLlvm(RangeArgs[1])
StepVal = self.HandleExprLlvm(RangeArgs[2])
else:
return None
if not EndVal:
return None
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="zip_end")
else:
EndVal = Gen.builder.ptrtoint(EndVal, ir.IntType(64), name="zip_end")
if not isinstance(StartVal.type, ir.IntType):
StartVal = ir.Constant(EndVal.type, 0)
if isinstance(StartVal.type, ir.IntType) and isinstance(EndVal.type, ir.IntType):
if StartVal.type.width < EndVal.type.width:
StartVal = Gen.builder.zext(StartVal, EndVal.type, name="zip_start_zext")
elif StartVal.type.width > EndVal.type.width:
EndVal = Gen.builder.zext(EndVal, StartVal.type, name="zip_end_zext")
if 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="zip_step_zext")
IdxVar: ir.AllocaInstr = Gen._allocaEntry(EndVal.type, name=f"zip_{target_name}_idx")
Gen._store(StartVal, IdxVar)
LoopVar: ir.AllocaInstr = Gen._allocaEntry(EndVal.type, name=target_name)
Gen.variables[target_name] = LoopVar
Gen._store(StartVal, LoopVar)
Gen._record_var_signedness(target_name, 'int')
return {
'iter_type': 'range',
'target_name': target_name,
'value_type': EndVal.type,
'idx_var': IdxVar,
'end_val': EndVal,
'step_val': StepVal,
'loop_var': LoopVar,
}
def _zip_setup_iter_class(self, arg: ast.expr, target_name: str, ClassName: str, Gen: "Translator.LlvmGen") -> dict | None:
IterVal: ir.Value | None = self.HandleExprLlvm(arg)
if not IterVal:
return None
return self._zip_setup_iter_class_from_val(IterVal, target_name, ClassName, Gen)
def _zip_setup_iter_class_from_val(self, IterVal: ir.Value, target_name: str, ClassName: str, Gen: "Translator.LlvmGen") -> dict | None:
Gen._UnregisterTempPtr(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"zip_cast_{ClassName}")
IterCall: ir.Function = Gen._get_function(f'{ClassName}.__iter__')
IterResult: ir.Value = Gen.builder.call(IterCall, [IterVal], name=f"zip_call_{ClassName}.__iter__")
Gen._UnregisterTempPtr(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"zip_cast_iter_{ClassName}")
NextCall: ir.Function = Gen._get_function(f'{ClassName}.__next__')
StopFlagPtr: ir.AllocaInstr = Gen._allocaEntry(ir.IntType(1), name=f"zip_stop_{target_name}")
Gen.builder.store(ir.Constant(ir.IntType(1), 0), StopFlagPtr)
NextReturnType: ir.Type = NextCall.function_type.return_type
LoopVar: ir.AllocaInstr = Gen._allocaEntry(NextReturnType, name=target_name)
Gen.variables[target_name] = LoopVar
return {
'iter_type': 'iter',
'target_name': target_name,
'value_type': NextReturnType,
'iter_obj': IterResult,
'next_call': NextCall,
'stop_flag': StopFlagPtr,
'loop_var': LoopVar,
}
def _zip_setup_str(self, StrVal: ir.Value, target_name: str, Gen: "Translator.LlvmGen") -> dict:
Gen._UnregisterTempPtr(StrVal)
PtrVal: ir.AllocaInstr = Gen.builder.alloca(ir.IntType(8).as_pointer(), name=f"zip_{target_name}_ptr")
Gen.builder.store(StrVal, PtrVal)
CharType: ir.IntType = ir.IntType(8)
LoopVar: ir.AllocaInstr = Gen._allocaEntry(CharType, name=target_name)
Gen.variables[target_name] = LoopVar
Gen._store(ir.Constant(CharType, 0), LoopVar)
Gen._record_var_signedness(target_name, 'char')
return {
'iter_type': 'str',
'target_name': target_name,
'value_type': CharType,
'ptr_var': PtrVal,
'loop_var': LoopVar,
}
def _zip_setup_array(self, ArrPtr: ir.Value, target_name: str, Gen: "Translator.LlvmGen") -> dict:
ElemType: ir.Type = ArrPtr.type.pointee.element
ArrayCount: int = ArrPtr.type.pointee.count
IdxVar: ir.AllocaInstr = Gen._allocaEntry(ir.IntType(32), name=f"zip_{target_name}_idx")
Gen.builder.store(ir.Constant(ir.IntType(32), 0), IdxVar)
LoopVar: ir.AllocaInstr = Gen._allocaEntry(ElemType, name=target_name)
Gen.variables[target_name] = LoopVar
Gen._store(ir.Constant(ElemType, None), LoopVar)
Gen._record_var_signedness(target_name, 'ptr')
return {
'iter_type': 'array',
'target_name': target_name,
'value_type': ElemType,
'arr_ptr': ArrPtr,
'idx_var': IdxVar,
'count_val': ir.Constant(ir.IntType(32), ArrayCount),
'loop_var': LoopVar,
}
def _zip_setup_ptr(self, PtrVal: ir.Value, target_name: str, ElemType: ir.IntType, Gen: "Translator.LlvmGen") -> dict:
TypedPtr: ir.Value = PtrVal
if isinstance(PtrVal.type, ir.PointerType) and PtrVal.type.pointee is not ElemType:
TypedPtr = Gen.builder.bitcast(PtrVal, ir.PointerType(ElemType), name="zip_ptr_cast")
Gen._UnregisterTempPtr(TypedPtr)
PtrCopy: ir.AllocaInstr = Gen.builder.alloca(ir.PointerType(ElemType), name=f"zip_{target_name}_ptr")
Gen.builder.store(TypedPtr, PtrCopy)
LoopVar: ir.AllocaInstr = Gen._allocaEntry(ElemType, name=target_name)
Gen.variables[target_name] = LoopVar
Gen._store(ir.Constant(ElemType, 0), LoopVar)
Gen._record_var_signedness(target_name, 'int')
return {
'iter_type': 'ptr',
'target_name': target_name,
'value_type': ElemType,
'ptr_var': PtrCopy,
'loop_var': LoopVar,
}
def _zip_setup_int(self, IterVal: ir.Value, target_name: str, Gen: "Translator.LlvmGen") -> dict:
EndVal: ir.Value = IterVal
if not isinstance(EndVal.type, ir.IntType):
EndVal = Gen.builder.ptrtoint(EndVal, ir.IntType(64), name="zip_int_end")
StartVal: ir.Value = ir.Constant(EndVal.type, 0)
StepVal: ir.Value = ir.Constant(EndVal.type, 1)
IdxVar: ir.AllocaInstr = Gen._allocaEntry(EndVal.type, name=f"zip_{target_name}_idx")
Gen._store(StartVal, IdxVar)
LoopVar: ir.AllocaInstr = Gen._allocaEntry(EndVal.type, name=target_name)
Gen.variables[target_name] = LoopVar
Gen._store(StartVal, LoopVar)
Gen._record_var_signedness(target_name, 'int')
return {
'iter_type': 'range',
'target_name': target_name,
'value_type': EndVal.type,
'idx_var': IdxVar,
'end_val': EndVal,
'step_val': StepVal,
'loop_var': LoopVar,
}
def _zip_get_next(self, info: dict, Gen: "Translator.LlvmGen") -> tuple[ir.Value, ir.Value]:
"""生成获取下一个值和停止条件的 IR。返回 (value, stopped_cond)。"""
iter_type: str = info['iter_type']
if iter_type == 'range':
idx: ir.Value = Gen._load(info['idx_var'], name="zip_range_idx")
stopped: ir.Value = Gen.builder.icmp_signed('>=', idx, info['end_val'], name="zip_range_stop")
next_idx: ir.Value = Gen.builder.add(idx, info['step_val'], name="zip_range_next")
Gen._store(next_idx, info['idx_var'])
return idx, stopped
if iter_type == 'iter':
Gen.builder.store(ir.Constant(ir.IntType(1), 0), info['stop_flag'])
args: list[ir.Value] = [info['iter_obj'], info['stop_flag']]
NextCall: ir.Function = info['next_call']
if NextCall and len(NextCall.function_type.args) > len(args):
last_param_type: ir.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: ir.AllocaInstr = Gen._allocaEntry(ir.PointerType(ir.IntType(8)), name="zip_eh_msg")
args.append(null_msg)
val: ir.Value = Gen.builder.call(NextCall, args, name="zip_iter_next")
stopped = Gen.builder.load(info['stop_flag'], name="zip_iter_stop")
return val, stopped
if iter_type == 'str':
ptr: ir.Value = Gen._load(info['ptr_var'], name="zip_str_ptr")
val = Gen._load(ptr, name="zip_str_val")
stopped = Gen.builder.icmp_signed('==', val, ir.Constant(ir.IntType(8), 0), name="zip_str_stop")
next_ptr: ir.Value = Gen.builder.gep(ptr, [ir.Constant(ir.IntType(32), 1)], name="zip_str_next")
Gen._store(next_ptr, info['ptr_var'])
return val, stopped
if iter_type == 'ptr':
ptr = Gen._load(info['ptr_var'], name="zip_ptr_cur")
val = Gen._load(ptr, name="zip_ptr_val")
elem_type: ir.IntType = info['value_type']
stopped = Gen.builder.icmp_signed('==', val, ir.Constant(elem_type, 0), name="zip_ptr_stop")
next_ptr = Gen.builder.gep(ptr, [ir.Constant(ir.IntType(32), 1)], name="zip_ptr_next")
Gen._store(next_ptr, info['ptr_var'])
return val, stopped
if iter_type == 'array':
idx = Gen._load(info['idx_var'], name="zip_arr_idx")
stopped = Gen.builder.icmp_signed('>=', idx, info['count_val'], name="zip_arr_stop")
elem_ptr: ir.Value = Gen.builder.gep(info['arr_ptr'], [ir.Constant(ir.IntType(32), 0), idx], name="zip_arr_elem")
val = Gen._load(elem_ptr, name="zip_arr_val")
next_idx = Gen.builder.add(idx, ir.Constant(ir.IntType(32), 1), name="zip_arr_next")
Gen._store(next_idx, info['idx_var'])
return val, stopped
# 回退:立即停止
return ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(1), 1)
# ===== enumerate() 内置支持 =====
def _HandleForEnumerateLlvm(self, Node: ast.For) -> None:
"""for idx, val in enumerate(iterable): — 带索引迭代
等价于 zip(range(0, ∞), iterable),但无需 len()。
idx 从 0 开始,每次迭代 +1。
"""
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
if not isinstance(Node.target, ast.Tuple):
return
Targets: list[ast.Name] = [t for t in Node.target.elts if isinstance(t, ast.Name)]
if len(Targets) != 2:
return
IdxTargetName: str = Targets[0].id
ValTargetName: str = Targets[1].id
IterArg: ast.expr = Node.iter.args[0]
# 设置内部可迭代对象(复用 zip 基础设施)
if ValTargetName in Gen._reg_values:
del Gen._reg_values[ValTargetName]
if IdxTargetName in Gen._reg_values:
del Gen._reg_values[IdxTargetName]
IterInfo: dict | None = self._zip_setup_iter(IterArg, ValTargetName, Gen)
if IterInfo is None:
return
# 索引计数器i32从 0 开始)
IdxType: ir.IntType = ir.IntType(32)
IdxVar: ir.AllocaInstr = Gen._allocaEntry(IdxType, name=f"enum_{IdxTargetName}")
Gen._store(ir.Constant(IdxType, 0), IdxVar)
IdxLoopVar: ir.AllocaInstr = Gen._allocaEntry(IdxType, name=IdxTargetName)
Gen.variables[IdxTargetName] = IdxLoopVar
Gen._store(ir.Constant(IdxType, 0), IdxLoopVar)
Gen._record_var_signedness(IdxTargetName, 'int')
CondBB: ir.Block = Gen.func.append_basic_block(name="enum.cond")
BodyBB: ir.Block = Gen.func.append_basic_block(name="enum.body")
StepBB: ir.Block = Gen.func.append_basic_block(name="enum.step")
ElseBB: ir.Block | None = Gen.func.append_basic_block(name="enum.else") if Node.orelse else None
AfterBB: ir.Block = Gen.func.append_basic_block(name="enum.end")
Gen.loop_break_targets.append(AfterBB)
Gen.loop_continue_targets.append(StepBB)
Gen.builder.branch(CondBB)
# Cond 块:获取下一个值和停止条件
Gen.builder.position_at_start(CondBB)
CurIdx: ir.Value = Gen._load(IdxVar, name="enum_idx")
Val, Stopped = self._zip_get_next(IterInfo, Gen)
IterInfo['last_value'] = Val
Gen.builder.cbranch(Stopped, ElseBB if ElseBB else AfterBB, BodyBB)
# Body 块:赋值并执行循环体
Gen.builder.position_at_start(BodyBB)
Gen._store(CurIdx, IdxLoopVar)
Gen._store(Val, IterInfo['loop_var'])
Gen._reg_values[IdxTargetName] = CurIdx
Gen._reg_values[ValTargetName] = Val
self.HandleBodyLlvm(Node.body)
if IdxTargetName in Gen._reg_values and Gen._reg_values[IdxTargetName] is CurIdx:
del Gen._reg_values[IdxTargetName]
if ValTargetName in Gen._reg_values and Gen._reg_values[ValTargetName] is Val:
del Gen._reg_values[ValTargetName]
if not Gen.builder.block.is_terminated:
Gen.builder.branch(StepBB)
# Step 块:索引 +1跳回 Cond
Gen.builder.position_at_start(StepBB)
CurIdx2: ir.Value = Gen._load(IdxVar, name="enum_idx2")
NextIdx: ir.Value = Gen.builder.add(CurIdx2, ir.Constant(IdxType, 1), name="enum_next")
Gen._store(NextIdx, IdxVar)
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)