382 lines
22 KiB
Python
382 lines
22 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 MatchHandle(BaseHandle):
|
||
def _HandleMatchLlvm(self, Node: ast.Match) -> None:
|
||
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
|
||
SubjectVal: ir.Value | None = self.HandleExprLlvm(Node.subject)
|
||
if not SubjectVal:
|
||
return
|
||
|
||
IsRenumMatch: bool = False
|
||
RenumName: str | None = None
|
||
SubjectPtr: ir.Value | None = None
|
||
if isinstance(Node.subject, ast.Name):
|
||
VarName: str = Node.subject.id
|
||
TypeInfo: "SymbolTable.SymbolInfo | None" = self.Trans.SymbolTable.lookup(VarName)
|
||
if TypeInfo and TypeInfo.IsRenum:
|
||
IsRenumMatch = True
|
||
RenumName = TypeInfo.Name
|
||
SubjectPtr = Gen._Load_var(VarName)
|
||
|
||
if not IsRenumMatch:
|
||
for case in Node.cases:
|
||
if isinstance(case.pattern, ast.MatchClass):
|
||
cls_node: ast.expr = case.pattern.cls
|
||
VariantName: str | None = None
|
||
QualifiedName: str | None = None
|
||
if isinstance(cls_node, ast.Name):
|
||
VariantName = cls_node.id
|
||
elif isinstance(cls_node, ast.Attribute):
|
||
VariantName = cls_node.attr
|
||
# Extract enum name for qualified lookup to avoid collision
|
||
# with same-name factory functions (e.g. def Ptr vs LLVMType.Ptr)
|
||
if isinstance(cls_node.value, ast.Attribute):
|
||
QualifiedName = f"{cls_node.value.attr}.{VariantName}"
|
||
elif isinstance(cls_node.value, ast.Name):
|
||
QualifiedName = f"{cls_node.value.id}.{VariantName}"
|
||
if VariantName:
|
||
# Try qualified name first (e.g., "LLVMType.Ptr") to avoid
|
||
# collision with same-name functions (e.g., def Ptr(...))
|
||
SymInfo: "SymbolTable.SymbolInfo" = None
|
||
if QualifiedName:
|
||
SymInfo = self.Trans.SymbolTable.lookup(QualifiedName)
|
||
if not (SymInfo and SymInfo.IsEnumMember):
|
||
SymInfo = self.Trans.SymbolTable.lookup(VariantName)
|
||
if SymInfo and SymInfo.IsEnumMember and SymInfo.EnumName:
|
||
EnumName: str = SymInfo.EnumName
|
||
EnumInfo: "SymbolTable.SymbolInfo" = self.Trans.SymbolTable.lookup(EnumName)
|
||
if EnumInfo and EnumInfo.IsRenum:
|
||
IsRenumMatch = True
|
||
RenumName = EnumName
|
||
if SubjectPtr is None:
|
||
SubjectPtr = self.HandleExprLlvm(Node.subject)
|
||
break
|
||
|
||
if IsRenumMatch and SubjectPtr:
|
||
self._HandleRenumMatchLlvm(Node, RenumName, SubjectPtr)
|
||
return
|
||
|
||
if not isinstance(SubjectVal.type, ir.IntType):
|
||
try:
|
||
SubjectVal = Gen.builder.ptrtoint(SubjectVal, ir.IntType(64), name="match_subj")
|
||
SubjectVal = Gen.builder.trunc(SubjectVal, ir.IntType(32), name="match_subj_i32")
|
||
except Exception: # 回退:ptrtoint 失败时直接返回
|
||
return
|
||
SwitchIntType: ir.IntType = SubjectVal.type
|
||
DefaultBB: ir.Block = Gen.func.append_basic_block(name="match.default")
|
||
AfterBB: ir.Block = Gen.func.append_basic_block(name="match.end")
|
||
CaseBBs: list[ir.Block] = []
|
||
CaseValues: list[ir.Value] = []
|
||
HasDefault: bool = False
|
||
HasNoBreak: list[bool] = []
|
||
for i, case in enumerate(Node.cases):
|
||
pattern: ast.pattern = case.pattern
|
||
if isinstance(pattern, ast.MatchValue):
|
||
Val: ir.Value | None = self.HandleExprLlvm(pattern.value)
|
||
if Val:
|
||
CaseVal: ir.Value
|
||
if isinstance(Val.type, ir.IntType):
|
||
if Val.type != SwitchIntType:
|
||
if Val.type.width > SwitchIntType.width:
|
||
CaseVal = ir.Constant(SwitchIntType, Val.constant & ((1 << SwitchIntType.width) - 1))
|
||
else:
|
||
CaseVal = ir.Constant(SwitchIntType, Val.constant)
|
||
else:
|
||
CaseVal = Val
|
||
else:
|
||
try:
|
||
CaseVal = Gen.builder.ptrtoint(Val, SwitchIntType, name=f"case_val_{i}")
|
||
except Exception: # 回退:ptrtoint 失败时设默认值 0
|
||
CaseVal = ir.Constant(SwitchIntType, 0)
|
||
CaseValues.append(CaseVal)
|
||
CaseBBs.append(Gen.func.append_basic_block(name=f"match.case_{i}"))
|
||
elif isinstance(pattern, ast.MatchOr):
|
||
for j, SubPattern in enumerate(pattern.patterns):
|
||
if isinstance(SubPattern, ast.MatchValue):
|
||
Val: ir.Value | None = self.HandleExprLlvm(SubPattern.value)
|
||
if Val:
|
||
CaseVal: ir.Value
|
||
if isinstance(Val.type, ir.IntType):
|
||
if Val.type != SwitchIntType:
|
||
if Val.type.width > SwitchIntType.width:
|
||
CaseVal = ir.Constant(SwitchIntType, Val.constant & ((1 << SwitchIntType.width) - 1))
|
||
else:
|
||
CaseVal = ir.Constant(SwitchIntType, Val.constant)
|
||
else:
|
||
CaseVal = Val
|
||
else:
|
||
try:
|
||
CaseVal = Gen.builder.ptrtoint(Val, SwitchIntType, name=f"case_val_{i}_{j}")
|
||
except Exception: # 回退:ptrtoint 失败时设默认值 0
|
||
CaseVal = ir.Constant(SwitchIntType, 0)
|
||
CaseValues.append(CaseVal)
|
||
CaseBBs.append(Gen.func.append_basic_block(name=f"match.case_{i}_{j}"))
|
||
elif isinstance(pattern, ast.MatchSingleton):
|
||
if pattern.value is None:
|
||
CaseValues.append(ir.Constant(SwitchIntType, 0))
|
||
CaseBBs.append(Gen.func.append_basic_block(name=f"match.case_{i}"))
|
||
elif isinstance(pattern, ast.MatchAs):
|
||
if pattern.pattern is None:
|
||
HasDefault = True
|
||
CaseBBs.append(DefaultBB)
|
||
elif isinstance(pattern, ast.MatchSequence):
|
||
HasDefault = True
|
||
CaseBBs.append(DefaultBB)
|
||
def _HasNoBreak(stmts: list[ast.stmt]) -> bool:
|
||
for stmt in stmts:
|
||
if isinstance(stmt, ast.Expr) and isinstance(stmt.value, ast.Call):
|
||
if isinstance(stmt.value.func, ast.Attribute):
|
||
if (isinstance(stmt.value.func.value, ast.Name) and
|
||
stmt.value.func.value.id == 'c' and
|
||
stmt.value.func.attr == 'NoBreak'):
|
||
return True
|
||
if getattr(stmt, 'body', None) and isinstance(stmt.body, list):
|
||
if _HasNoBreak(stmt.body):
|
||
return True
|
||
if getattr(stmt, 'orelse', None):
|
||
if isinstance(stmt.orelse, list) and _HasNoBreak(stmt.orelse):
|
||
return True
|
||
return False
|
||
HasNoBreak.append(_HasNoBreak(case.body) if case.body else False)
|
||
if not HasDefault:
|
||
CaseBBs.append(DefaultBB)
|
||
SwitchCases: list[tuple[ir.Value, ir.Block]] = [(val, bb) for val, bb in zip(CaseValues, CaseBBs) if bb != DefaultBB]
|
||
switch_instr: ir.SwitchInstr = Gen.builder.switch(SubjectVal, DefaultBB)
|
||
for val, bb in SwitchCases:
|
||
switch_instr.add_case(val, bb)
|
||
CaseIdx: int = 0
|
||
for i, case in enumerate(Node.cases):
|
||
pattern: ast.pattern = case.pattern
|
||
if isinstance(pattern, ast.MatchOr):
|
||
NumSubCases: int = len(pattern.patterns)
|
||
for j in range(NumSubCases):
|
||
if CaseIdx < len(CaseBBs) and CaseBBs[CaseIdx] != DefaultBB:
|
||
Gen.builder.position_at_start(CaseBBs[CaseIdx])
|
||
self.HandleBodyLlvm(case.body)
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
elif isinstance(pattern, (ast.MatchValue, ast.MatchSingleton)):
|
||
if CaseIdx < len(CaseBBs) and CaseBBs[CaseIdx] != DefaultBB:
|
||
Gen.builder.position_at_start(CaseBBs[CaseIdx])
|
||
self.HandleBodyLlvm(case.body)
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
elif isinstance(pattern, ast.MatchAs):
|
||
if pattern.pattern is None:
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
self.HandleBodyLlvm(case.body)
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
elif isinstance(pattern, ast.MatchSequence):
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
self.HandleBodyLlvm(case.body)
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
if not HasDefault:
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
Gen.builder.branch(AfterBB)
|
||
elif not any(isinstance(c.pattern, ast.MatchAs) and c.pattern.pattern is None for c in Node.cases) and not any(isinstance(c.pattern, ast.MatchSequence) for c in Node.cases):
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
Gen.builder.branch(AfterBB)
|
||
Gen.builder.position_at_start(AfterBB)
|
||
|
||
def _HandleRenumMatchLlvm(self, Node: ast.Match, RenumName: str, SubjectPtr: ir.Value) -> None:
|
||
Gen: "Translator.LlvmGen" = self.Trans.LlvmGen
|
||
if isinstance(SubjectPtr.type, ir.PointerType) and isinstance(SubjectPtr.type.pointee, ir.PointerType):
|
||
SubjectPtr = Gen._load(SubjectPtr, name="Load_match_subj")
|
||
# Bug fix: HandleExprLlvm 可能对 REnum 字段执行了 load,返回值而非指针。
|
||
# 情况1: SubjectPtr 不是指针类型(如 IntType)→ alloca REnum 结构体 + store
|
||
# 情况2: SubjectPtr 是指针但 pointee 不是结构体(如 i32* 指向 __tag)→ bitcast
|
||
NeedAlloca: bool = not isinstance(SubjectPtr.type, ir.PointerType)
|
||
NeedBitcast: bool = (isinstance(SubjectPtr.type, ir.PointerType) and
|
||
not isinstance(SubjectPtr.type.pointee, (ir.IdentifiedStructType, ir.LiteralStructType)))
|
||
if NeedAlloca or NeedBitcast:
|
||
RenumStructType: Any = Gen.structs.get(RenumName)
|
||
if RenumStructType is None:
|
||
# 跨模块编译时 REnum 结构体可能尚未在 Gen.structs 中注册
|
||
# (如 import llvmlite 后 LLVMType 仅在符号表,Gen.structs 无条目)。
|
||
# 使用 _get_or_create_struct 按需创建(可能为 opaque)。
|
||
RenumStructType = Gen._get_or_create_struct(RenumName)
|
||
if NeedAlloca:
|
||
AllocaPtr: ir.Value = Gen._allocaEntry(RenumStructType, name="match_subj_alloca")
|
||
CastedPtr: ir.Value = Gen.builder.bitcast(AllocaPtr, ir.PointerType(SubjectPtr.type), name="match_subj_cast")
|
||
Gen._store(SubjectPtr, CastedPtr)
|
||
SubjectPtr = AllocaPtr
|
||
else:
|
||
SubjectPtr = Gen.builder.bitcast(SubjectPtr, ir.PointerType(RenumStructType), name="match_subj_recast")
|
||
# REnum 布局为 { i32 __tag, <payload> },tag 在偏移 0。
|
||
# 使用 bitcast 到 i32* 替代 gep [0,0],兼容 opaque 结构体
|
||
# (跨模块编译时 REnum 结构体可能未 set_body,gep 会失败)。
|
||
tag_ptr: ir.Value = Gen.builder.bitcast(SubjectPtr, ir.PointerType(ir.IntType(32)), name="match_tag_ptr")
|
||
TagVal: ir.Value = Gen._load(tag_ptr, name="match_tag_val")
|
||
DefaultBB: ir.Block = Gen.func.append_basic_block(name="match.default")
|
||
AfterBB: ir.Block = Gen.func.append_basic_block(name="match.end")
|
||
CaseBBs: list[ir.Block] = []
|
||
CaseValues: list[ir.Value] = []
|
||
CaseBindings: list[list[tuple[str, str, ir.Type, int]]] = []
|
||
HasDefault: bool = False
|
||
HasNoBreak: list[bool] = []
|
||
for i, case in enumerate(Node.cases):
|
||
pattern: ast.pattern = case.pattern
|
||
bindings: list[tuple[str, str, ir.Type, int]] = []
|
||
if isinstance(pattern, ast.MatchClass):
|
||
cls_node: ast.expr = pattern.cls
|
||
VariantName: str | None = None
|
||
QualifiedName: str | None = None
|
||
if isinstance(cls_node, ast.Name):
|
||
VariantName = cls_node.id
|
||
elif isinstance(cls_node, ast.Attribute):
|
||
VariantName = cls_node.attr
|
||
if isinstance(cls_node.value, ast.Attribute):
|
||
QualifiedName = f"{cls_node.value.attr}.{VariantName}"
|
||
elif isinstance(cls_node.value, ast.Name):
|
||
QualifiedName = f"{cls_node.value.id}.{VariantName}"
|
||
if VariantName:
|
||
TagValue: int | None = None
|
||
SymInfo: "SymbolTable.SymbolInfo" = None
|
||
if QualifiedName:
|
||
SymInfo = self.Trans.SymbolTable.lookup(QualifiedName)
|
||
if not (SymInfo and SymInfo.IsEnumMember):
|
||
SymInfo = self.Trans.SymbolTable.lookup(VariantName)
|
||
if SymInfo and SymInfo.IsEnumMember:
|
||
TagValue = SymInfo.value
|
||
if TagValue is not None:
|
||
CaseValues.append(ir.Constant(ir.IntType(32), TagValue))
|
||
CaseBB: ir.Block = Gen.func.append_basic_block(name=f"match.case_{VariantName}")
|
||
CaseBBs.append(CaseBB)
|
||
NestedStructName: str = f"{RenumName}_{VariantName}"
|
||
if NestedStructName in Gen.structs:
|
||
members: list[tuple[str, ir.Type]] = Gen.class_members.get(NestedStructName, [])
|
||
payLoad_members: list[tuple[str, ir.Type]] = [(n, t) for n, t in members if n != '__tag']
|
||
for j, sub_pat in enumerate(pattern.patterns):
|
||
if isinstance(sub_pat, ast.MatchAs) and sub_pat.name and j < len(payLoad_members):
|
||
bindings.append((sub_pat.name, payLoad_members[j][0], payLoad_members[j][1], j))
|
||
CaseBindings.append(bindings)
|
||
else:
|
||
CaseValues.append(ir.Constant(ir.IntType(32), 0))
|
||
CaseBBs.append(Gen.func.append_basic_block(name=f"match.case_{i}"))
|
||
CaseBindings.append([])
|
||
elif isinstance(pattern, ast.MatchValue):
|
||
Val: ir.Value | None = self.HandleExprLlvm(pattern.value)
|
||
if Val:
|
||
CaseVal: ir.Value
|
||
if isinstance(Val.type, ir.IntType):
|
||
CaseVal = Val
|
||
else:
|
||
try:
|
||
CaseVal = Gen.builder.ptrtoint(Val, ir.IntType(32), name=f"case_val_{i}")
|
||
except Exception: # 回退:ptrtoint 失败时设默认值 0
|
||
CaseVal = ir.Constant(ir.IntType(32), 0)
|
||
CaseValues.append(CaseVal)
|
||
CaseBBs.append(Gen.func.append_basic_block(name=f"match.case_{i}"))
|
||
CaseBindings.append([])
|
||
elif isinstance(pattern, ast.MatchAs):
|
||
if pattern.pattern is None:
|
||
HasDefault = True
|
||
CaseBBs.append(DefaultBB)
|
||
CaseBindings.append([])
|
||
elif isinstance(pattern, ast.MatchSequence):
|
||
HasDefault = True
|
||
CaseBBs.append(DefaultBB)
|
||
CaseBindings.append([])
|
||
def _HasNoBreak(stmts: list[ast.stmt]) -> bool:
|
||
for stmt in stmts:
|
||
if isinstance(stmt, ast.Expr) and isinstance(stmt.value, ast.Call):
|
||
if isinstance(stmt.value.func, ast.Attribute):
|
||
if (isinstance(stmt.value.func.value, ast.Name) and
|
||
stmt.value.func.value.id == 'c' and
|
||
stmt.value.func.attr == 'NoBreak'):
|
||
return True
|
||
if getattr(stmt, 'body', None) and isinstance(stmt.body, list):
|
||
if _HasNoBreak(stmt.body):
|
||
return True
|
||
if getattr(stmt, 'orelse', None):
|
||
if isinstance(stmt.orelse, list) and _HasNoBreak(stmt.orelse):
|
||
return True
|
||
return False
|
||
HasNoBreak.append(_HasNoBreak(case.body) if case.body else False)
|
||
if not HasDefault:
|
||
CaseBBs.append(DefaultBB)
|
||
SwitchCases: list[tuple[ir.Value, ir.Block]] = [(val, bb) for val, bb in zip(CaseValues, CaseBBs) if bb != DefaultBB]
|
||
switch_instr: ir.SwitchInstr = Gen.builder.switch(TagVal, DefaultBB)
|
||
for val, bb in SwitchCases:
|
||
switch_instr.add_case(val, bb)
|
||
CaseIdx: int = 0
|
||
for i, case in enumerate(Node.cases):
|
||
pattern: ast.pattern = case.pattern
|
||
if isinstance(pattern, ast.MatchClass):
|
||
if CaseIdx < len(CaseBBs) and CaseBBs[CaseIdx] != DefaultBB:
|
||
Gen.builder.position_at_start(CaseBBs[CaseIdx])
|
||
bindings: list[tuple[str, str, ir.Type, int]] = CaseBindings[CaseIdx] if CaseIdx < len(CaseBindings) else []
|
||
cls_node: ast.expr = pattern.cls
|
||
VariantName: str | None = None
|
||
if isinstance(cls_node, ast.Name):
|
||
VariantName = cls_node.id
|
||
elif isinstance(cls_node, ast.Attribute):
|
||
VariantName = cls_node.attr
|
||
if VariantName:
|
||
NestedStructName: str = f"{RenumName}_{VariantName}"
|
||
if NestedStructName in Gen.structs:
|
||
NestedStructType: ir.Type = Gen.structs[NestedStructName]
|
||
NestedStructPtrType: ir.PointerType = ir.PointerType(NestedStructType)
|
||
variant_ptr: ir.Value = Gen.builder.bitcast(SubjectPtr, NestedStructPtrType, name=f"match_cast_{VariantName}")
|
||
for bind_name, member_name, member_type, member_idx in bindings:
|
||
elem_ptr: ir.Value = Gen.builder.gep(variant_ptr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), member_idx + 1)], name=f"match_{bind_name}")
|
||
# REnum 嵌套结构体共享 max_variant_struct 布局,成员类型可能与
|
||
# 结构体字段类型不一致。bitcast 到成员类型以确保后续 load 得到正确类型。
|
||
if isinstance(elem_ptr.type, ir.PointerType) and elem_ptr.type.pointee != member_type:
|
||
elem_ptr = Gen.builder.bitcast(elem_ptr, ir.PointerType(member_type), name=f"match_{bind_name}_cast")
|
||
Gen.variables[bind_name] = elem_ptr
|
||
self.HandleBodyLlvm(case.body)
|
||
for bind_name, _, _, _ in bindings:
|
||
if bind_name in Gen.variables:
|
||
del Gen.variables[bind_name]
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
elif isinstance(pattern, ast.MatchValue):
|
||
if CaseIdx < len(CaseBBs) and CaseBBs[CaseIdx] != DefaultBB:
|
||
Gen.builder.position_at_start(CaseBBs[CaseIdx])
|
||
self.HandleBodyLlvm(case.body)
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
elif isinstance(pattern, ast.MatchAs):
|
||
if pattern.pattern is None:
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
self.HandleBodyLlvm(case.body)
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
elif isinstance(pattern, ast.MatchSequence):
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
self.HandleBodyLlvm(case.body)
|
||
if not Gen.builder.block.is_terminated:
|
||
if not HasNoBreak[i]:
|
||
Gen.builder.branch(AfterBB)
|
||
CaseIdx += 1
|
||
if not HasDefault:
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
Gen.builder.branch(AfterBB)
|
||
elif not any(isinstance(c.pattern, ast.MatchAs) and c.pattern.pattern is None for c in Node.cases) and not any(isinstance(c.pattern, ast.MatchSequence) for c in Node.cases):
|
||
Gen.builder.position_at_start(DefaultBB)
|
||
Gen.builder.branch(AfterBB)
|
||
Gen.builder.position_at_start(AfterBB) |