1073 lines
58 KiB
Python
1073 lines
58 KiB
Python
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) |