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, },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)