可用的回归测试通过的标准版本
This commit is contained in:
@@ -1,333 +1,333 @@
|
||||
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):
|
||||
Gen = self.Trans.LlvmGen
|
||||
SubjectVal = self.HandleExprLlvm(Node.subject)
|
||||
if not SubjectVal:
|
||||
return
|
||||
|
||||
IsRenumMatch = False
|
||||
RenumName = None
|
||||
SubjectPtr = None
|
||||
if isinstance(Node.subject, ast.Name):
|
||||
VarName = Node.subject.id
|
||||
if VarName in self.Trans.SymbolTable:
|
||||
TypeInfo = self.Trans.SymbolTable[VarName]
|
||||
if getattr(TypeInfo, 'IsRenum', False):
|
||||
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 = case.pattern.cls
|
||||
VariantName = None
|
||||
if isinstance(cls_node, ast.Name):
|
||||
VariantName = cls_node.id
|
||||
elif isinstance(cls_node, ast.Attribute):
|
||||
VariantName = cls_node.attr
|
||||
if VariantName and VariantName in self.Trans.SymbolTable:
|
||||
SymInfo = self.Trans.SymbolTable[VariantName]
|
||||
if getattr(SymInfo, 'IsEnumMember', False) and getattr(SymInfo, 'EnumName', None):
|
||||
EnumName = SymInfo.EnumName
|
||||
if EnumName in self.Trans.SymbolTable:
|
||||
EnumInfo = self.Trans.SymbolTable[EnumName]
|
||||
if getattr(EnumInfo, 'IsRenum', False):
|
||||
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 = SubjectVal.type
|
||||
DefaultBB = Gen.func.append_basic_block(name="match.default")
|
||||
AfterBB = Gen.func.append_basic_block(name="match.end")
|
||||
CaseBBs = []
|
||||
CaseValues = []
|
||||
HasDefault = False
|
||||
HasNoBreak = []
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
if isinstance(pattern, ast.MatchValue):
|
||||
Val = self.HandleExprLlvm(pattern.value)
|
||||
if Val:
|
||||
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 = self.HandleExprLlvm(SubPattern.value)
|
||||
if Val:
|
||||
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):
|
||||
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 = [(val, bb) for val, bb in zip(CaseValues, CaseBBs) if bb != DefaultBB]
|
||||
switch_instr = Gen.builder.switch(SubjectVal, DefaultBB)
|
||||
for val, bb in SwitchCases:
|
||||
switch_instr.add_case(val, bb)
|
||||
CaseIdx = 0
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
if isinstance(pattern, ast.MatchOr):
|
||||
NumSubCases = 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, RenumName, SubjectPtr):
|
||||
Gen = self.Trans.LlvmGen
|
||||
if isinstance(SubjectPtr.type, ir.PointerType) and isinstance(SubjectPtr.type.pointee, ir.PointerType):
|
||||
SubjectPtr = Gen._load(SubjectPtr, name="load_match_subj")
|
||||
tag_ptr = Gen.builder.gep(SubjectPtr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), 0)], name="match_tag_ptr")
|
||||
TagVal = Gen._load(tag_ptr, name="match_tag_val")
|
||||
DefaultBB = Gen.func.append_basic_block(name="match.default")
|
||||
AfterBB = Gen.func.append_basic_block(name="match.end")
|
||||
CaseBBs = []
|
||||
CaseValues = []
|
||||
CaseBindings = []
|
||||
HasDefault = False
|
||||
HasNoBreak = []
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
bindings = []
|
||||
if isinstance(pattern, ast.MatchClass):
|
||||
cls_node = pattern.cls
|
||||
VariantName = None
|
||||
if isinstance(cls_node, ast.Name):
|
||||
VariantName = cls_node.id
|
||||
elif isinstance(cls_node, ast.Attribute):
|
||||
VariantName = cls_node.attr
|
||||
if VariantName:
|
||||
TagValue = None
|
||||
if VariantName in self.Trans.SymbolTable:
|
||||
SymInfo = self.Trans.SymbolTable[VariantName]
|
||||
if getattr(SymInfo, 'IsEnumMember', False):
|
||||
TagValue = SymInfo.value
|
||||
if TagValue is not None:
|
||||
CaseValues.append(ir.Constant(ir.IntType(32), TagValue))
|
||||
CaseBB = Gen.func.append_basic_block(name=f"match.case_{VariantName}")
|
||||
CaseBBs.append(CaseBB)
|
||||
NestedStructName = f"{RenumName}_{VariantName}"
|
||||
if NestedStructName in Gen.structs:
|
||||
members = Gen.class_members.get(NestedStructName, [])
|
||||
payload_members = [(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 = self.HandleExprLlvm(pattern.value)
|
||||
if Val:
|
||||
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):
|
||||
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 = [(val, bb) for val, bb in zip(CaseValues, CaseBBs) if bb != DefaultBB]
|
||||
switch_instr = Gen.builder.switch(TagVal, DefaultBB)
|
||||
for val, bb in SwitchCases:
|
||||
switch_instr.add_case(val, bb)
|
||||
CaseIdx = 0
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
if isinstance(pattern, ast.MatchClass):
|
||||
if CaseIdx < len(CaseBBs) and CaseBBs[CaseIdx] != DefaultBB:
|
||||
Gen.builder.position_at_start(CaseBBs[CaseIdx])
|
||||
bindings = CaseBindings[CaseIdx] if CaseIdx < len(CaseBindings) else []
|
||||
cls_node = pattern.cls
|
||||
VariantName = None
|
||||
if isinstance(cls_node, ast.Name):
|
||||
VariantName = cls_node.id
|
||||
elif isinstance(cls_node, ast.Attribute):
|
||||
VariantName = cls_node.attr
|
||||
if VariantName:
|
||||
NestedStructName = f"{RenumName}_{VariantName}"
|
||||
if NestedStructName in Gen.structs:
|
||||
NestedStructType = Gen.structs[NestedStructName]
|
||||
NestedStructPtrType = ir.PointerType(NestedStructType)
|
||||
variant_ptr = Gen.builder.bitcast(SubjectPtr, NestedStructPtrType, name=f"match_cast_{VariantName}")
|
||||
for bind_name, member_name, member_type, member_idx in bindings:
|
||||
elem_ptr = Gen.builder.gep(variant_ptr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), member_idx + 1)], name=f"match_{bind_name}")
|
||||
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)
|
||||
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):
|
||||
Gen = self.Trans.LlvmGen
|
||||
SubjectVal = self.HandleExprLlvm(Node.subject)
|
||||
if not SubjectVal:
|
||||
return
|
||||
|
||||
IsRenumMatch = False
|
||||
RenumName = None
|
||||
SubjectPtr = None
|
||||
if isinstance(Node.subject, ast.Name):
|
||||
VarName = Node.subject.id
|
||||
if VarName in self.Trans.SymbolTable:
|
||||
TypeInfo = self.Trans.SymbolTable[VarName]
|
||||
if getattr(TypeInfo, 'IsRenum', False):
|
||||
IsRenumMatch = True
|
||||
RenumName = TypeInfo.Name
|
||||
SubjectPtr = Gen._loadVar(VarName)
|
||||
|
||||
if not IsRenumMatch:
|
||||
for case in Node.cases:
|
||||
if isinstance(case.pattern, ast.MatchClass):
|
||||
cls_node = case.pattern.cls
|
||||
VariantName = None
|
||||
if isinstance(cls_node, ast.Name):
|
||||
VariantName = cls_node.id
|
||||
elif isinstance(cls_node, ast.Attribute):
|
||||
VariantName = cls_node.attr
|
||||
if VariantName and VariantName in self.Trans.SymbolTable:
|
||||
SymInfo = self.Trans.SymbolTable[VariantName]
|
||||
if getattr(SymInfo, 'IsEnumMember', False) and getattr(SymInfo, 'EnumName', None):
|
||||
EnumName = SymInfo.EnumName
|
||||
if EnumName in self.Trans.SymbolTable:
|
||||
EnumInfo = self.Trans.SymbolTable[EnumName]
|
||||
if getattr(EnumInfo, 'IsRenum', False):
|
||||
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 = SubjectVal.type
|
||||
DefaultBB = Gen.func.append_basic_block(name="match.default")
|
||||
AfterBB = Gen.func.append_basic_block(name="match.end")
|
||||
CaseBBs = []
|
||||
CaseValues = []
|
||||
HasDefault = False
|
||||
HasNoBreak = []
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
if isinstance(pattern, ast.MatchValue):
|
||||
Val = self.HandleExprLlvm(pattern.value)
|
||||
if Val:
|
||||
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 = self.HandleExprLlvm(SubPattern.value)
|
||||
if Val:
|
||||
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):
|
||||
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 = [(val, bb) for val, bb in zip(CaseValues, CaseBBs) if bb != DefaultBB]
|
||||
switch_instr = Gen.builder.switch(SubjectVal, DefaultBB)
|
||||
for val, bb in SwitchCases:
|
||||
switch_instr.add_case(val, bb)
|
||||
CaseIdx = 0
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
if isinstance(pattern, ast.MatchOr):
|
||||
NumSubCases = 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, RenumName, SubjectPtr):
|
||||
Gen = self.Trans.LlvmGen
|
||||
if isinstance(SubjectPtr.type, ir.PointerType) and isinstance(SubjectPtr.type.pointee, ir.PointerType):
|
||||
SubjectPtr = Gen._load(SubjectPtr, name="Load_match_subj")
|
||||
tag_ptr = Gen.builder.gep(SubjectPtr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), 0)], name="match_tag_ptr")
|
||||
TagVal = Gen._load(tag_ptr, name="match_tag_val")
|
||||
DefaultBB = Gen.func.append_basic_block(name="match.default")
|
||||
AfterBB = Gen.func.append_basic_block(name="match.end")
|
||||
CaseBBs = []
|
||||
CaseValues = []
|
||||
CaseBindings = []
|
||||
HasDefault = False
|
||||
HasNoBreak = []
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
bindings = []
|
||||
if isinstance(pattern, ast.MatchClass):
|
||||
cls_node = pattern.cls
|
||||
VariantName = None
|
||||
if isinstance(cls_node, ast.Name):
|
||||
VariantName = cls_node.id
|
||||
elif isinstance(cls_node, ast.Attribute):
|
||||
VariantName = cls_node.attr
|
||||
if VariantName:
|
||||
TagValue = None
|
||||
if VariantName in self.Trans.SymbolTable:
|
||||
SymInfo = self.Trans.SymbolTable[VariantName]
|
||||
if getattr(SymInfo, 'IsEnumMember', False):
|
||||
TagValue = SymInfo.value
|
||||
if TagValue is not None:
|
||||
CaseValues.append(ir.Constant(ir.IntType(32), TagValue))
|
||||
CaseBB = Gen.func.append_basic_block(name=f"match.case_{VariantName}")
|
||||
CaseBBs.append(CaseBB)
|
||||
NestedStructName = f"{RenumName}_{VariantName}"
|
||||
if NestedStructName in Gen.structs:
|
||||
members = Gen.class_members.get(NestedStructName, [])
|
||||
payLoad_members = [(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 = self.HandleExprLlvm(pattern.value)
|
||||
if Val:
|
||||
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):
|
||||
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 = [(val, bb) for val, bb in zip(CaseValues, CaseBBs) if bb != DefaultBB]
|
||||
switch_instr = Gen.builder.switch(TagVal, DefaultBB)
|
||||
for val, bb in SwitchCases:
|
||||
switch_instr.add_case(val, bb)
|
||||
CaseIdx = 0
|
||||
for i, case in enumerate(Node.cases):
|
||||
pattern = case.pattern
|
||||
if isinstance(pattern, ast.MatchClass):
|
||||
if CaseIdx < len(CaseBBs) and CaseBBs[CaseIdx] != DefaultBB:
|
||||
Gen.builder.position_at_start(CaseBBs[CaseIdx])
|
||||
bindings = CaseBindings[CaseIdx] if CaseIdx < len(CaseBindings) else []
|
||||
cls_node = pattern.cls
|
||||
VariantName = None
|
||||
if isinstance(cls_node, ast.Name):
|
||||
VariantName = cls_node.id
|
||||
elif isinstance(cls_node, ast.Attribute):
|
||||
VariantName = cls_node.attr
|
||||
if VariantName:
|
||||
NestedStructName = f"{RenumName}_{VariantName}"
|
||||
if NestedStructName in Gen.structs:
|
||||
NestedStructType = Gen.structs[NestedStructName]
|
||||
NestedStructPtrType = ir.PointerType(NestedStructType)
|
||||
variant_ptr = Gen.builder.bitcast(SubjectPtr, NestedStructPtrType, name=f"match_cast_{VariantName}")
|
||||
for bind_name, member_name, member_type, member_idx in bindings:
|
||||
elem_ptr = Gen.builder.gep(variant_ptr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), member_idx + 1)], name=f"match_{bind_name}")
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user