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

368 lines
21 KiB
Python
Raw 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 Any
import logging
import llvmlite.ir as ir
from lib.constants.config import mode as _config_mode
class ExprGenMixin:
def emit_constant(self, value: Any, type_name: str = 'int') -> ir.Value:
if type_name == 'int':
return ir.Constant(ir.IntType(32), int(value))
elif type_name == 'uint64_t' or type_name == 'unsigned long long':
return ir.Constant(ir.IntType(64), int(value))
elif type_name == 'uint32_t' or type_name == 'unsigned int':
return ir.Constant(ir.IntType(32), int(value))
elif type_name == 'uint16_t' or type_name == 'unsigned short':
return ir.Constant(ir.IntType(16), int(value))
elif type_name == 'uint8_t' or type_name == 'unsigned char':
return ir.Constant(ir.IntType(8), int(value))
elif type_name == 'int64_t' or type_name == 'long long':
return ir.Constant(ir.IntType(64), int(value))
elif type_name == 'int32_t' or type_name == 'int':
return ir.Constant(ir.IntType(32), int(value))
elif type_name == 'int16_t' or type_name == 'short':
return ir.Constant(ir.IntType(16), int(value))
elif type_name == 'int8_t' or type_name == 'char':
return ir.Constant(ir.IntType(8), int(value))
elif type_name == 'float' or type_name == 'float32_t' or type_name == 'FLOAT32':
return ir.Constant(ir.FloatType(), float(value))
elif type_name == 'double' or type_name == 'float64_t' or type_name == 'FLOAT64':
return ir.Constant(ir.DoubleType(), float(value))
elif type_name == 'float16_t' or type_name == 'FLOAT16':
return ir.Constant(ir.IntType(16), int(value))
elif type_name == 'float8_t' or type_name == 'FLOAT8':
return ir.Constant(ir.IntType(8), int(value))
elif type_name == 'float128_t' or type_name == 'FLOAT128':
return ir.Constant(ir.IntType(128), int(value))
elif type_name == 'string':
encoded: bytes = value.encode('utf-8') + b'\x00'
str_type: ir.ArrayType = ir.ArrayType(ir.IntType(8), len(encoded))
str_const: ir.Constant = ir.Constant(str_type, bytearray(encoded))
global_var: ir.GlobalVariable = ir.GlobalVariable(self.module, str_type, name=f"str_const_{self.string_const_counter}")
self.string_const_counter += 1
global_var.initializer = str_const
global_var.linkage = 'internal'
return self.builder.bitcast(global_var, ir.PointerType(ir.IntType(8)), name="str")
return ir.Constant(ir.IntType(32), int(value))
def emit_binary_op(self, op: str, left: ir.Value, right: ir.Value, is_unsigned: bool = False) -> ir.Value:
if isinstance(left.type, ir.PointerType) and isinstance(right.type, ir.IntType):
# 数组指针需要先 decay 为元素指针,否则 GEP 会以整个数组为步长
# 例如 [64 x i8]* + 1 会前进 64 字节而非 1 字节
if isinstance(left.type.pointee, ir.ArrayType):
zero: ir.Constant = ir.Constant(ir.IntType(32), 0)
left = self.builder.gep(left, [zero, zero], name="array_decay")
right = self.builder.zext(right, ir.IntType(64), name="int2ptrint")
if op == '+' or op == 'Add':
return self.builder.gep(left, [right], name="ptr_add")
elif op == '-' or op == 'Sub':
return self.builder.gep(left, [self.builder.neg(right, name="neg")], name="ptr_sub")
if isinstance(left.type, ir.IntType) and isinstance(right.type, ir.PointerType):
if isinstance(right.type.pointee, ir.ArrayType):
zero: ir.Constant = ir.Constant(ir.IntType(32), 0)
right = self.builder.gep(right, [zero, zero], name="array_decay_rev")
left = self.builder.zext(left, ir.IntType(64), name="int2ptrint2")
if op == '+' or op == 'Add':
return self.builder.gep(right, [left], name="ptr_add_rev")
is_float: bool = isinstance(left.type, (ir.FloatType, ir.DoubleType)) or \
isinstance(right.type, (ir.FloatType, ir.DoubleType))
if is_float:
ftype: ir.Type
if isinstance(left.type, (ir.FloatType, ir.DoubleType)):
ftype = left.type
elif isinstance(right.type, (ir.FloatType, ir.DoubleType)):
ftype = right.type
else:
ftype = ir.DoubleType()
if not isinstance(left.type, (ir.FloatType, ir.DoubleType)):
left = self._to_float(left, ftype)
elif left.type != ftype:
if isinstance(left.type, ir.FloatType) and isinstance(ftype, ir.DoubleType):
left = self.builder.fpext(left, ftype, name="fpext_l")
else:
left = self.builder.fptrunc(left, ftype, name="fptrunc_l")
if not isinstance(right.type, (ir.FloatType, ir.DoubleType)):
right = self._to_float(right, ftype)
elif right.type != ftype:
if isinstance(right.type, ir.FloatType) and isinstance(ftype, ir.DoubleType):
right = self.builder.fpext(right, ftype, name="fpext_r")
else:
right = self.builder.fptrunc(right, ftype, name="fptrunc_r")
float_ops: dict[str, Any] = {
'+': self.builder.fadd, 'Add': self.builder.fadd,
'-': self.builder.fsub, 'Sub': self.builder.fsub,
'*': self.builder.fmul, 'Mult': self.builder.fmul,
'/': self.builder.fdiv, 'Div': self.builder.fdiv,
'%': self.builder.frem, 'Mod': self.builder.frem,
}
float_cmp_ops: dict[str, str] = {
'==': '==', 'Eq': '==', '!=': '!=', 'NotEq': '!=',
'<': '<', 'Lt': '<', '<=': '<=', 'LtE': '<=',
'>': '>', 'Gt': '>', '>=': '>=', 'GtE': '>=',
}
if op in float_ops:
return float_ops[op](left, right, name=op)
if op in float_cmp_ops:
return self.builder.fcmp_ordered(float_cmp_ops[op], left, right, name=op)
ops: dict[str, Any]
if is_unsigned:
ops = {
'+': self.builder.add, 'Add': self.builder.add,
'-': self.builder.sub, 'Sub': self.builder.sub,
'*': self.builder.mul, 'Mult': self.builder.mul,
'/': self.builder.udiv, 'Div': self.builder.udiv,
'%': self.builder.urem, 'Mod': self.builder.urem,
'//': self.builder.udiv, 'FloorDiv': self.builder.udiv,
}
else:
ops = {
'+': self.builder.add, 'Add': self.builder.add,
'-': self.builder.sub, 'Sub': self.builder.sub,
'*': self.builder.mul, 'Mult': self.builder.mul,
'/': self.builder.sdiv, 'Div': self.builder.sdiv,
'%': self.builder.srem, 'Mod': self.builder.srem,
'//': self.builder.sdiv, 'FloorDiv': self.builder.sdiv,
}
cmp_ops: dict[str, str] = {
'==': '==', 'Eq': '==', '!=': '!=', 'NotEq': '!=',
'<': '<', 'Lt': '<', '<=': '<=', 'LtE': '<=',
'>': '>', 'Gt': '>', '>=': '>=', 'GtE': '>=',
}
logic_ops: dict[str, Any] = {
'&&': self.builder.and_, 'And': self.builder.and_,
'||': self.builder.or_, 'Or': self.builder.or_,
}
if op in ('>>', 'RShift') and not is_unsigned:
if isinstance(left.type, ir.IntType):
if left.type.width == 8:
is_unsigned = True
elif left.type.width == 16:
is_unsigned = True
elif left.type.width == 32:
if self._last_var_name:
is_unsigned = self._is_var_unsigned(self._last_var_name)
elif left.type.width == 64:
if self._last_var_name:
is_unsigned = self._is_var_unsigned(self._last_var_name)
shift_ops: dict[str, Any] = {
'>>': self.builder.lshr if is_unsigned else self.builder.ashr,
'RShift': self.builder.lshr if is_unsigned else self.builder.ashr,
'<<': self.builder.shl, 'LShift': self.builder.shl,
}
bit_ops: dict[str, Any] = {
'&': self.builder.and_,
'|': self.builder.or_,
'^': self.builder.xor,
}
if op == '**' or op == 'Pow':
if isinstance(left, ir.Constant) and isinstance(right, ir.Constant):
try:
lv: Any = left.constant
rv: Any = right.constant
if isinstance(lv, int) and isinstance(rv, int) and rv >= 0:
result: int = lv ** rv
return ir.Constant(left.type, result)
except Exception as _e:
if _config_mode == "strict":
logging.warning(f"异常被忽略: {_e}")
pass
ftype: ir.Type
if isinstance(left.type, (ir.FloatType, ir.DoubleType)):
ftype = left.type
elif isinstance(right.type, (ir.FloatType, ir.DoubleType)):
ftype = right.type
else:
ftype = ir.DoubleType()
lf: ir.Value = self._to_float(left, ftype)
rf: ir.Value = self._to_float(right, ftype)
powf: ir.Function | None = self.module.globals.get('llvm.pow.f64')
if not powf:
fnty: ir.FunctionType = ir.FunctionType(ir.DoubleType(), [ir.DoubleType(), ir.DoubleType()])
powf = ir.Function(self.module, fnty, name='llvm.pow.f64')
if lf.type != ir.DoubleType():
lf = self.builder.fpext(lf, ir.DoubleType(), name="ext_f64")
if rf.type != ir.DoubleType():
rf = self.builder.fpext(rf, ir.DoubleType(), name="ext_f64")
result: ir.Value = self.builder.call(powf, [lf, rf], name="pow")
both_int: bool = isinstance(left.type, ir.IntType) and isinstance(right.type, ir.IntType)
right_is_pos_int_const: bool = isinstance(right, ir.Constant) and isinstance(right.constant, int) and right.constant >= 0
if both_int and right_is_pos_int_const:
return self.builder.fptosi(result, left.type, name="pow2int")
if isinstance(left.type, ir.FloatType):
return self.builder.fptrunc(result, ir.FloatType(), name="pow2f32")
return result
if op in ops:
try:
return ops[op](left, right, name=op)
except Exception as e:
raise ValueError(f"{e}{self._get_node_info()}")
if op in shift_ops:
try:
return shift_ops[op](left, right, name=op)
except Exception as e:
raise ValueError(f"{e}{self._get_node_info()}")
if op in bit_ops:
try:
return bit_ops[op](left, right, name=op)
except Exception as e:
raise ValueError(f"{e}{self._get_node_info()}")
if op in cmp_ops:
try:
if is_unsigned:
return self.builder.icmp_unsigned(cmp_ops[op], left, right, name=op)
return self.builder.icmp_signed(cmp_ops[op], left, right, name=op)
except Exception as e:
raise ValueError(f"{e}{self._get_node_info()}")
if op in logic_ops:
try:
return logic_ops[op](left, right, name=op)
except Exception as e:
raise ValueError(f"{e}{self._get_node_info()}")
return left
def _to_float(self, val: ir.Value, ftype: ir.Type | None = None) -> ir.Value:
if ftype is None:
ftype = ir.DoubleType()
if isinstance(val.type, (ir.FloatType, ir.DoubleType)):
if val.type == ftype:
return val
if isinstance(val.type, ir.FloatType) and isinstance(ftype, ir.DoubleType):
return self.builder.fpext(val, ftype, name="fpext")
return self.builder.fptrunc(val, ftype, name="fptrunc")
if isinstance(val.type, ir.IntType):
if val.type.width == 1:
val = self.builder.zext(val, ir.IntType(32), name="bool2i32")
return self.builder.sitofp(val, ftype, name="int2float")
return val
def emit_return(self, value: ir.Value | None = None) -> None:
if not self.builder or self.builder.block.is_terminated:
return
if value is not None:
self._unregister_local_heap_ptr(value)
self._emit_local_heap_frees()
if self._variadic_info and self._variadic_info.get('va_start_called'):
va_list_ptr: ir.Value = self._variadic_info['va_list_ptr']
self.emit_va_end(va_list_ptr)
if value is None:
if self.func and hasattr(self.func, 'ftype'):
ret_type: ir.Type = self.func.ftype.return_type
if isinstance(ret_type, ir.IntType):
self.builder.ret(ir.Constant(ret_type, 0))
elif isinstance(ret_type, ir.PointerType):
self.builder.ret(ir.Constant(ret_type, None))
elif isinstance(ret_type, (ir.FloatType, ir.DoubleType)):
self.builder.ret(ir.Constant(ret_type, 0.0))
elif isinstance(ret_type, ir.BaseStructType):
if ret_type.elements:
zero_val: ir.Constant = ir.Constant(ret_type, [ir.Constant(et, None) if isinstance(et, (ir.PointerType, ir.IdentifiedStructType, ir.LiteralStructType, ir.ArrayType)) else ir.Constant(et, 0) if isinstance(et, (ir.IntType, ir.FloatType, ir.DoubleType)) else ir.Constant(et, ir.Undefined) for et in ret_type.elements])
else:
zero_val = ir.Constant(ret_type, None)
self.builder.ret(zero_val)
elif isinstance(ret_type, ir.ArrayType):
self.builder.ret(ir.Constant(ret_type, None))
else:
self.builder.ret_void()
else:
self.builder.ret_void()
else:
if self.func and hasattr(self.func, 'ftype'):
ret_type: ir.Type = self.func.ftype.return_type
if ret_type != value.type:
if isinstance(ret_type, ir.PointerType) and isinstance(value.type, ir.PointerType):
value = self.builder.bitcast(value, ret_type, name="ret_cast")
elif isinstance(value.type, ir.PointerType) and not isinstance(ret_type, ir.PointerType):
if isinstance(value.type.pointee, ir.IntType) and isinstance(ret_type, ir.IntType) and value.type.pointee.width == ret_type.width:
value = self._load(value, name="ret_Load")
elif isinstance(ret_type, ir.IntType):
value = self.builder.ptrtoint(value, ret_type, name="ret_ptrtoint")
elif isinstance(value.type.pointee, ret_type.__class__) and not isinstance(value.type.pointee, ir.IntType):
value = self._load(value, name="ret_Load")
elif isinstance(value.type.pointee, ir.PointerType):
Loaded: ir.Value = self._load(value, name="ret_deref")
if Loaded.type == ret_type:
value = Loaded
elif isinstance(ret_type, ir.IntType):
value = self.builder.ptrtoint(Loaded, ret_type, name="ret_ptrtoint")
elif isinstance(ret_type, ir.IntType) and isinstance(value.type, ir.IntType):
if value.type.width < ret_type.width:
if value.type.width == 1:
# i1 是布尔值(来自 icmp/fcmp/__contains__无符号语义必须 zext
value = self.builder.zext(value, ret_type, name="ret_zext_bool")
else:
is_unsigned: bool = id(value) in self._unsigned_results
if is_unsigned:
value = self.builder.zext(value, ret_type, name="ret_zext")
else:
value = self.builder.sext(value, ret_type, name="ret_sext")
elif value.type.width > ret_type.width:
value = self.builder.trunc(value, ret_type, name="ret_trunc")
elif isinstance(ret_type, ir.PointerType) and isinstance(value.type, ir.IntType):
value = self.builder.inttoptr(value, ret_type, name="ret_inttoptr")
elif isinstance(ret_type, (ir.FloatType, ir.DoubleType)) and isinstance(value.type, ir.IntType):
value = self.builder.sitofp(value, ret_type, name="ret_int2float")
elif isinstance(ret_type, ir.IntType) and isinstance(value.type, (ir.FloatType, ir.DoubleType)):
value = self.builder.fptosi(value, ret_type, name="ret_float2int")
elif isinstance(ret_type, (ir.FloatType, ir.DoubleType)) and isinstance(value.type, (ir.FloatType, ir.DoubleType)):
if ret_type != value.type:
if isinstance(value.type, ir.FloatType) and isinstance(ret_type, ir.DoubleType):
value = self.builder.fpext(value, ret_type, name="ret_fpext")
elif isinstance(value.type, ir.DoubleType) and isinstance(ret_type, ir.FloatType):
value = self.builder.fptrunc(value, ret_type, name="ret_fptrunc")
self.builder.ret(value)
def _emit_printf(self, args: list[ir.Value], unsigned_flags: list[bool] | None = None, sep: str = " ", end: str = "\n") -> ir.Value | None:
if not args:
return None
if unsigned_flags is None:
unsigned_flags = []
format_str: str = ""
cast_args: list[ir.Value] = []
for i, arg in enumerate(args):
is_u: bool = i < len(unsigned_flags) and unsigned_flags[i]
if isinstance(arg.type, ir.IntType):
if arg.type.width == 1:
# i1 是布尔值(来自 icmp/fcmp/__contains__zext 到 i32 后用 %d
format_str += "%d"
ext: ir.Value = self.builder.zext(arg, ir.IntType(32), name="bool2int")
cast_args.append(ext)
elif arg.type.width == 8:
format_str += "%c"
ext: ir.Value = self.builder.zext(arg, ir.IntType(32), name="char2int")
cast_args.append(ext)
elif arg.type.width == 64:
format_str += "%llu" if is_u else "%lld"
cast_args.append(arg)
else:
format_str += "%u" if is_u else "%d"
cast_args.append(arg)
elif isinstance(arg.type, (ir.FloatType, ir.DoubleType)):
format_str += "%f"
cast_args.append(arg)
elif isinstance(arg.type, ir.PointerType):
if isinstance(arg.type.pointee, ir.IntType) and arg.type.pointee.width == 8:
format_str += "%s"
cast_args.append(arg)
else:
format_str += "%d"
int_val: ir.Value = self.builder.ptrtoint(arg, ir.IntType(64), name="ptr2int")
trunc_val: ir.Value = self.builder.trunc(int_val, ir.IntType(32), name="trunc")
cast_args.append(trunc_val)
else:
format_str += "%d"
cast_args.append(arg)
if sep is not None and i < len(args) - 1:
format_str += sep
if end is not None:
format_str += end
string_type: ir.ArrayType = ir.ArrayType(ir.IntType(8), len(format_str.encode('utf-8')) + 1)
string_const: ir.Constant = ir.Constant(string_type, bytearray(format_str + '\x00', 'utf-8'))
global_var: ir.GlobalVariable = ir.GlobalVariable(self.module, string_type, name=f"str_const_{self.string_const_counter}")
self.string_const_counter += 1
global_var.initializer = string_const
global_var.linkage = 'internal'
format_ptr: ir.Value = self.builder.bitcast(global_var, ir.PointerType(ir.IntType(8)))
call_args: list[ir.Value] = [format_ptr] + cast_args
printf_func: ir.Function | None = self.get_or_declare_c_func('printf', ir.FunctionType(ir.IntType(32), [ir.PointerType(ir.IntType(8))], var_arg=True))
return self.builder.call(printf_func, call_args, name="call_printf")