Files
TransPyC/lib/core/Handles/HandlesMatch.py
2026-07-18 19:25:40 +08:00

382 lines
22 KiB
Python
Raw Permalink 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 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_bodygep 会失败)。
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)