from __future__ import annotations from typing import TYPE_CHECKING, Any import ast import re import llvmlite.ir as ir if TYPE_CHECKING: from lib.core.translator import Translator from lib.core.LlvmCodeGenerator import LlvmCodeGenerator from lib.core.Handles.HandlesBase import BaseHandle from lib.core.SymbolUtils import IsTModule from lib.includes import t # 猴子补丁 llvmlite.ir.InlineAsm 以支持 asm_dialect 参数 # 添加保护避免重复补丁 _InlineAsmPatched = getattr(ir.InlineAsm, '_transpyc_patched', False) if not _InlineAsmPatched: _OrigInlineAsmInit = ir.InlineAsm.__init__ _OrigInlineAsmDescr = ir.InlineAsm.descr def _patched_inline_asm_init(self, ftype, asm, constraint, side_effect=False, asm_dialect=None): _OrigInlineAsmInit(self, ftype, asm, constraint, side_effect) self.asm_dialect = asm_dialect def _patched_inline_asm_descr(self, buf): sideeffect = 'sideeffect' if self.side_effect else '' dialect_str = getattr(self, 'asm_dialect', None) if dialect_str == 'intel': dialect = 'inteldialect' else: dialect = '' fmt = 'asm {sideeffect} {dialect} "{asm}", "{constraint}"' buf.append(fmt.format(sideeffect=sideeffect, dialect=dialect, asm=self.asm, constraint=self.constraint)) ir.InlineAsm.__init__ = _patched_inline_asm_init ir.InlineAsm.descr = _patched_inline_asm_descr ir.InlineAsm._transpyc_patched = True class ExprAsmHandle(BaseHandle): @staticmethod def _format_to_dialect(asm_format: str | None) -> str | None: """将用户指定的 format 参数转换为 LLVM asm_dialect 值""" fmt: str = asm_format.lower() if asm_format else 'intel' if fmt == 'att': return 'att' elif fmt == 'arm': return None # ARM 没有专门的 dialect 标记 else: return 'intel' # 默认 Intel def _HandleAsmLlvm(self, Node: ast.Call) -> ir.Value | None: Gen: LlvmCodeGenerator = self.Trans.LlvmGen op_arg: Any = None input_arg: Any = None output_arg: Any = None asm_format: str = 'Intel' # 汇编格式: Intel / AT&T / ARM for kw in Node.keywords: if kw.arg == 'op': op_arg = kw.value elif kw.arg in ('input', 'inp', 'inputs'): input_arg = kw.value elif kw.arg in ('output', 'out', 'outputs'): output_arg = kw.value elif kw.arg == 'format': if isinstance(kw.value, ast.Constant) and isinstance(kw.value.value, str): asm_format = kw.value.value if len(Node.args) >= 1: pos_op: Any = None if len(Node.args) >= 2 and op_arg is None: pos_op = Node.args[1] return self._HandleAsmWithOpLlvm(Node.args[0], op_arg or pos_op, input_arg, output_arg, asm_format) if len(Node.args) < 2: return ir.Constant(ir.IntType(32), 1) args: list[ast.AST] = Node.args if len(args) == 4: output_type_arg: ast.AST = args[0] asm_template_arg: ast.AST = args[1] constraint_arg: ast.AST = args[2] clobber_arg: ast.AST = args[3] output_type: ir.Type | None = self._infer_expr_llvm_type(output_type_arg) if not output_type: output_type = ir.IntType(32) asm_template: str = '' if isinstance(asm_template_arg, ast.Constant) and isinstance(asm_template_arg.value, str): asm_template = self._prepare_intel_asm(self._mangle_asm_calls(asm_template_arg.value)) constraint: str = '' if isinstance(constraint_arg, ast.Constant) and isinstance(constraint_arg.value, str): constraint = constraint_arg.value clobber: str = '' if isinstance(clobber_arg, ast.Constant) and isinstance(clobber_arg.value, str): clobber = clobber_arg.value full_constraint: str = constraint func_type: ir.FunctionType = ir.FunctionType(output_type, []) asm_dialect: str | None = self._format_to_dialect(asm_format) inline_asm: ir.InlineAsm = ir.InlineAsm(func_type, asm_template, constraint, side_effect=True, asm_dialect=asm_dialect) result: Any = Gen.builder.call(inline_asm, [], name="asm_result") return result elif len(args) >= 2: output_type_arg: ast.AST = args[0] asm_template_arg: ast.AST = args[1] output_type: ir.Type | None = self._infer_expr_llvm_type(output_type_arg) if not output_type: output_type = ir.IntType(32) asm_template: str = '' if isinstance(asm_template_arg, ast.Constant) and isinstance(asm_template_arg.value, str): asm_template = self._prepare_intel_asm(self._mangle_asm_calls(asm_template_arg.value)) output_constraints: list[str] = [] output_values: list[Any] = [] if len(args) > 2 and isinstance(args[2], (ast.List, ast.Tuple)): for item in args[2].elts: if isinstance(item, (ast.List, ast.Tuple)) and len(item.elts) >= 2: constraint_node: ast.AST = item.elts[0] value_node: ast.AST = item.elts[1] if isinstance(constraint_node, ast.Constant) and isinstance(constraint_node.value, str): output_constraints.append(constraint_node.value) value: Any = self.HandleExprLlvm(value_node) if value: output_values.append(value) input_constraints: list[str] = [] input_values: list[Any] = [] if len(args) > 3 and isinstance(args[3], (ast.List, ast.Tuple)): for item in args[3].elts: if isinstance(item, (ast.List, ast.Tuple)) and len(item.elts) >= 2: constraint_node: ast.AST = item.elts[0] value_node: ast.AST = item.elts[1] if isinstance(constraint_node, ast.Constant) and isinstance(constraint_node.value, str): input_constraints.append(constraint_node.value) value: Any = self.HandleExprLlvm(value_node) if value: input_values.append(value) clobbers: list[str] = [] if len(args) > 4 and isinstance(args[4], (ast.List, ast.Tuple)): for item in args[4].elts: if isinstance(item, ast.Constant) and isinstance(item.value, str): clobbers.append(item.value) constraint_parts: list[str] = [] for oc in output_constraints: constraint_parts.append(oc) input_start_idx: int = len(output_constraints) for ic in input_constraints: constraint_parts.append(ic) full_constraint: str = ",".join(constraint_parts) if clobbers: if full_constraint: full_constraint += "," for c in clobbers: full_constraint += f"~{{{c}}}" operands: list[Any] = output_values + input_values param_types: list[ir.Type] = [v.type for v in operands] func_type: ir.FunctionType = ir.FunctionType(output_type, param_types) asm_dialect: str | None = self._format_to_dialect(asm_format) inline_asm: ir.InlineAsm = ir.InlineAsm(func_type, asm_template, full_constraint, side_effect=False, asm_dialect=asm_dialect) result: Any = Gen.builder.call(inline_asm, operands, name="asm_result") return result return None def _prepare_intel_asm(self, template: str) -> str: template = re.sub(r'%(\d+)', r'$\1', template) template = re.sub(r'%([a-zA-Z][a-zA-Z0-9]*)', r'\1', template) return template def _normalize_asm_template(self, template: str) -> str: lines = template.split('\n') if len(lines) <= 1: return template first_line = lines[0] rest_lines = lines[1:] stripped = [] for line in rest_lines: s = line.lstrip() if s: stripped.append(s) else: stripped.append('') return '\n'.join([first_line] + stripped) def _process_joined_str_template(self, node: ast.AST, out_input_ops: list[Any], out_input_constraints: list[str], out_output_ops: list[Any], out_output_constraints: list[str]) -> str: if not isinstance(node, ast.JoinedStr): if isinstance(node, ast.Constant) and isinstance(node.value, str): return node.value return '' raw_parts: list[tuple[str, str | int]] = [] for value in node.values: if isinstance(value, ast.Constant) and isinstance(value.value, str): raw_parts.append(('text', value.value)) elif isinstance(value, ast.FormattedValue): expr: ast.AST = value.value if isinstance(expr, ast.Call) and isinstance(expr.func, ast.Attribute): if isinstance(expr.func.value, ast.Name) and expr.func.value.id == 'c': if expr.func.attr == 'AsmInp': if expr.args: v: Any = self.HandleExprLlvm(expr.args[0]) if v: c: str = 'r' if len(expr.args) >= 2: c = self._parse_asm_descr(expr.args[1]) if not c: c = 'r' out_input_ops.append(v) out_input_constraints.append(c) raw_parts.append(('input', len(out_input_ops) - 1)) elif expr.func.attr == 'AsmOut': if expr.args: target_ptr: Any = self._get_var_ptr_for_asm(expr.args[0]) if target_ptr: v = target_ptr else: v = self.HandleExprLlvm(expr.args[0]) if v: c: str = '=r' if len(expr.args) >= 2: c = self._parse_asm_descr(expr.args[1]) if not c: c = '=r' if not c.startswith('='): c = '=' + c out_output_ops.append(v) out_output_constraints.append(c) raw_parts.append(('output', len(out_output_ops) - 1)) else: fallback: Any = self.HandleExprLlvm(expr) if fallback and isinstance(fallback.type, ir.IntType): out_input_ops.append(fallback) out_input_constraints.append('r') raw_parts.append(('input', len(out_input_ops) - 1)) else: raw_parts.append(('text', '')) num_outputs: int = len(out_output_ops) parts: list[str] = [] for kind, data in raw_parts: if kind == 'text': parts.append(data) elif kind == 'output': parts.append(f'%{data}') elif kind == 'input': parts.append(f'%{num_outputs + data}') return ''.join(parts) _X86_REGISTERS = frozenset({ 'rax', 'rbx', 'rcx', 'rdx', 'rsi', 'rdi', 'rbp', 'rsp', 'r8', 'r9', 'r10', 'r11', 'r12', 'r13', 'r14', 'r15', 'eax', 'ebx', 'ecx', 'edx', 'esi', 'edi', 'ebp', 'esp', 'ax', 'bx', 'cx', 'dx', 'si', 'di', 'bp', 'sp', 'al', 'bl', 'cl', 'dl', 'ah', 'bh', 'ch', 'dh', 'sil', 'dil', 'bpl', 'spl', 'r8d', 'r9d', 'r10d', 'r11d', 'r12d', 'r13d', 'r14d', 'r15d', 'r8w', 'r9w', 'r10w', 'r11w', 'r12w', 'r13w', 'r14w', 'r15w', 'r8b', 'r9b', 'r10b', 'r11b', 'r12b', 'r13b', 'r14b', 'r15b', }) def _mangle_asm_calls(self, template: str) -> str: Gen: LlvmCodeGenerator = self.Trans.LlvmGen sha1: str = getattr(Gen, 'module_sha1', '') if not sha1: return template def replace_call(m: re.Match[str]) -> str: prefix: str = m.group(1) func_name: str = m.group(2) if func_name.lower() in self._X86_REGISTERS: return m.group(0) mangled: str = Gen._mangle_func_name(func_name) if '.' in mangled: return f'{prefix}\\22{mangled}\\22' return f"{prefix}{mangled}" return re.sub(r'(call\s+)(\w+)', replace_call, template) def _HandleAsmWithOpLlvm(self, asm_template_arg: ast.AST, op_arg: ast.AST | None = None, input_arg: ast.AST | None = None, output_arg: ast.AST | None = None, asm_format: str = 'Intel') -> ir.Value: Gen: LlvmCodeGenerator = self.Trans.LlvmGen input_ops: list[Any] = [] input_constraints: list[str] = [] output_ops: list[Any] = [] output_constraints: list[str] = [] if output_arg and isinstance(output_arg, (ast.List, ast.Tuple)): for item in output_arg.elts: if isinstance(item, ast.Call) and isinstance(item.func, ast.Attribute): if item.func.value.id == 'c' and item.func.attr == 'AsmOut': if len(item.args) >= 1: target_ptr: Any = self._get_var_ptr_for_asm(item.args[0]) if target_ptr: output_ops.append(target_ptr) else: value: Any = self.HandleExprLlvm(item.args[0]) if value: output_ops.append(value) constraint: str = '=r' if len(item.args) >= 2: base_constraint: str = self._parse_asm_descr(item.args[1]) if base_constraint: if base_constraint.startswith('='): constraint = base_constraint else: constraint = '=' + base_constraint output_constraints.append(constraint) raw_template: str if isinstance(asm_template_arg, ast.JoinedStr): raw_template = self._normalize_asm_template( self._process_joined_str_template( asm_template_arg, input_ops, input_constraints, output_ops, output_constraints ) ) elif isinstance(asm_template_arg, ast.Constant) and isinstance(asm_template_arg.value, str): raw_template = self._normalize_asm_template(asm_template_arg.value) else: raw_template = '' raw_template = self._mangle_asm_calls(raw_template) asm_dialect: str | None = self._format_to_dialect(asm_format) asm_template: str if asm_format == 'AT&T': asm_template = raw_template # AT&T 不需要转换 elif asm_format == 'ARM': asm_template = raw_template # ARM 不需要转换 else: asm_template = self._prepare_intel_asm(raw_template) # Intel → LLVM IR 格式 if input_arg and isinstance(input_arg, (ast.List, ast.Tuple)): for item in input_arg.elts: if isinstance(item, ast.Call) and isinstance(item.func, ast.Attribute): if item.func.value.id == 'c' and item.func.attr == 'AsmInp': if len(item.args) >= 1: value: Any = self.HandleExprLlvm(item.args[0]) if value: input_ops.append(value) constraint: str = 'r' if len(item.args) >= 2: constraint = self._parse_asm_descr(item.args[1]) if not constraint: constraint = 'r' input_constraints.append(constraint) elif isinstance(item, (ast.List, ast.Tuple)) and len(item.elts) >= 2: value_node: ast.AST = item.elts[0] descr_node: ast.AST = item.elts[1] value: Any = self.HandleExprLlvm(value_node) if value: input_ops.append(value) constraint: str = self._parse_asm_descr(descr_node) if not constraint: constraint = 'r' input_constraints.append(constraint) clobbers: list[str] = [] if op_arg and isinstance(op_arg, (ast.List, ast.Tuple)): for item in op_arg.elts: clobber: str = self._parse_asm_descr(item) if clobber: clobbers.append(clobber) constraint_parts: list[str] = [] for oc in output_constraints: c: str = oc if oc.startswith('=') else ('=' + oc) constraint_parts.append(c) for ic in input_constraints: if ic in ('d', 'edx', 'rdx'): constraint_parts.append('{dx}') else: constraint_parts.append(ic) full_constraint: str = ",".join(constraint_parts) constraint_to_reg: dict[str, str] = { 'a': 'rax', 'b': 'rbx', 'c': 'rcx', 'd': 'rdx', 'S': 'rsi', 'D': 'rdi', 'A': 'rax', 'U': 'r8', } used_regs: set[str] = set() for ic in input_constraints: if ic in constraint_to_reg: used_regs.add(constraint_to_reg[ic]) for oc in output_constraints: base_oc: str = oc.lstrip('=').lstrip('+') if base_oc in constraint_to_reg: used_regs.add(constraint_to_reg[base_oc]) if clobbers: has_memory: bool = 'memory' in clobbers if has_memory and output_ops: clobbers = [c for c in clobbers if c != 'memory'] clobber_final: list[str] if asm_format == 'ARM': # ARM 不需要 x86 的 e→r 寄存器名转换 clobber_final = [c for c in clobbers if c not in used_regs] else: clobber_64bit: list[str] = [] for c in clobbers: if c in ['eax', 'ebx', 'ecx', 'edx', 'esi', 'edi', 'ebp', 'esp']: clobber_64bit.append(c.replace('e', 'r')) else: clobber_64bit.append(c) clobber_final = [c for c in clobber_64bit if c not in used_regs] if clobber_final: if full_constraint: full_constraint += "," full_constraint += ",".join([f"~{{{c}}}" for c in clobber_final]) if output_ops: single_out: bool = len(output_ops) == 1 if single_out: op: Any = output_ops[0] ret_type: ir.Type if isinstance(op, ir.AllocaInstr): ret_type = op.type.pointee elif isinstance(op, ir.GlobalVariable): ret_type = op.type.pointee elif isinstance(op, (ir.GEPInstr, ir.CastInstr)): ret_type = op.type.pointee if isinstance(op.type, ir.PointerType) else op.type else: ret_type = op.type else: def _get_ret_type(op: Any) -> ir.Type: if isinstance(op, (ir.AllocaInstr, ir.GlobalVariable)): return op.type.pointee return op.type ret_type = ir.LiteralStructType([_get_ret_type(op) for op in output_ops]) if isinstance(ret_type, ir.PointerType) and isinstance(ret_type.pointee, ir.IdentifiedStructType): ret_type = ir.IntType(64) operands: list[Any] = input_ops param_types: list[ir.Type] = [op.type for op in operands] func_type: ir.FunctionType = ir.FunctionType(ret_type, param_types) inline_asm: ir.InlineAsm = ir.InlineAsm(func_type, asm_template, full_constraint, side_effect=len(output_ops) == 0, asm_dialect=asm_dialect) result: Any = Gen.builder.call(inline_asm, operands, name="asm_result") if single_out: target: Any = output_ops[0] if isinstance(target, ir.AllocaInstr): if target.type.pointee != result.type: i64t: ir.IntType = ir.IntType(64) iv: Any = Gen.builder.ptrtoint(result, i64t) ptr: Any = Gen.builder.bitcast(target, ir.PointerType(i64t)) Gen.builder.store(iv, ptr) else: Gen.builder.store(result, target) elif isinstance(target, (ir.GlobalVariable, ir.GEPInstr, ir.CastInstr)): target_pointee: ir.Type = target.type.pointee if target_pointee != result.type: if isinstance(target_pointee, ir.PointerType) and isinstance(result, ir.Instruction) and result.type == ir.IntType(64): ptr_val: Any = Gen.builder.inttoptr(result, target_pointee) Gen.builder.store(ptr_val, target) else: i64t: ir.IntType = ir.IntType(64) iv: Any = Gen.builder.ptrtoint(result, i64t) ptr: Any = Gen.builder.bitcast(target, ir.PointerType(i64t)) Gen.builder.store(iv, ptr) else: Gen.builder.store(result, target) else: dst: Any = Gen._allocaEntry(result.type) Gen.builder.store(result, dst) else: for i, out_op in enumerate(output_ops): val: Any = Gen.builder.extract_value(result, i, name=f"asm_out_{i}") target: Any = out_op if isinstance(out_op, ir.AllocaInstr) else Gen._allocaEntry(val.type) if isinstance(target, ir.AllocaInstr) and target.type.pointee != val.type: i64t: ir.IntType = ir.IntType(64) iv: Any = Gen.builder.ptrtoint(val, i64t) ptr: Any = Gen.builder.bitcast(target, ir.PointerType(i64t)) Gen.builder.store(iv, ptr) else: Gen.builder.store(val, target) else: operands: list[Any] = input_ops param_types: list[ir.Type] = [op.type for op in operands] func_type: ir.FunctionType = ir.FunctionType(ir.VoidType(), param_types) inline_asm: ir.InlineAsm = ir.InlineAsm(func_type, asm_template, full_constraint, side_effect=True, asm_dialect=asm_dialect) Gen.builder.call(inline_asm, operands, name="asm_call") return ir.Constant(ir.IntType(32), 1) def _get_var_ptr_for_asm(self, node: ast.AST) -> ir.Value | None: Gen: LlvmCodeGenerator = self.Trans.LlvmGen if isinstance(node, ast.Name): var_name: str = node.id if var_name in Gen.variables and Gen.variables[var_name] is not None: ptr: Any = Gen.variables[var_name] if isinstance(ptr, ir.AllocaInstr): return ptr if isinstance(ptr, ir.GlobalVariable): return ptr elif isinstance(node, ast.Attribute): obj_ptr: Any = self._get_var_ptr_for_asm(node.value) if obj_ptr is not None: obj_type: ir.Type = obj_ptr.type.pointee if isinstance(obj_type, ir.PointerType) and isinstance(obj_type.pointee, ir.IdentifiedStructType): Loaded: Any = Gen.builder.load(obj_ptr, name="asm_attr_Load") struct_type: ir.IdentifiedStructType = obj_type.pointee class_name: str | None = Gen._resolve_class_for_var(node.value.id) if isinstance(node.value, ast.Name) else None if class_name: offset: int = Gen._get_member_offset(node.attr, class_name) gep: Any = Gen.builder.gep(Loaded, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), offset)], inbounds=True) return gep for idx, elem_type in enumerate(struct_type.elements): gep: Any = Gen.builder.gep(Loaded, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), idx)], inbounds=True) return gep elif isinstance(obj_type, ir.IdentifiedStructType): class_name: str | None = Gen._resolve_class_for_var(node.value.id) if isinstance(node.value, ast.Name) else None if class_name: offset: int = Gen._get_member_offset(node.attr, class_name) gep: Any = Gen.builder.gep(obj_ptr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), offset)], inbounds=True) return gep for idx, elem_type in enumerate(obj_type.elements): gep: Any = Gen.builder.gep(obj_ptr, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), idx)], inbounds=True) return gep return None def _parse_asm_descr(self, expr: ast.AST) -> str: if isinstance(expr, ast.BinOp): left_val: str = self._parse_asm_descr(expr.left) right_val: str = self._parse_asm_descr(expr.right) return left_val + right_val elif isinstance(expr, ast.Attribute): if isinstance(expr.value, ast.Attribute): if (isinstance(expr.value.value, ast.Name) and IsTModule(expr.value.value.id, self.Trans.SymbolTable) and expr.value.attr == 'ASM_DESCR'): AttrName: str = expr.attr return getattr(t.ASM_DESCR, AttrName, "") elif isinstance(expr.value, ast.Name) and IsTModule(expr.value.id, self.Trans.SymbolTable): AttrName: str = expr.attr return getattr(t.ASM_DESCR, AttrName, "") elif isinstance(expr, ast.Constant): return expr.value return "" def _infer_expr_llvm_type(self, expr: ast.AST) -> ir.Type | None: if isinstance(expr, ast.Name): type_name: str = expr.id if hasattr(t, type_name): ctype: Any = getattr(t, type_name) if isinstance(ctype, type) and issubclass(ctype, t.CType): size: Any = getattr(ctype, 'Size', None) if size is not None: return ir.IntType(size) if type_name == 'CFloat': return ir.FloatType() elif type_name == 'CDouble': return ir.DoubleType() elif isinstance(expr, ast.Attribute): if isinstance(expr.value, ast.Name) and IsTModule(expr.value.id, self.Trans.SymbolTable): return self._infer_expr_llvm_type(ast.Name(id=expr.attr, ctx=ast.Load())) return None