271 lines
13 KiB
Python
271 lines
13 KiB
Python
from __future__ import annotations
|
||
from typing import TYPE_CHECKING, Any
|
||
if TYPE_CHECKING:
|
||
from lib.core.translator import Translator
|
||
from lib.core.LlvmCodeGenerator import LlvmCodeGenerator
|
||
from lib.core.Handles.HandlesBase import BaseHandle
|
||
import ast
|
||
import re
|
||
import llvmlite.ir as ir
|
||
|
||
|
||
class ExprFormatHandle(BaseHandle):
|
||
def _HandleJoinedStrLlvm(self, Node: ast.JoinedStr) -> ir.Value:
|
||
Gen: LlvmCodeGenerator = self.Trans.LlvmGen
|
||
if not Gen or not Gen.builder:
|
||
return ir.Constant(ir.IntType(8).as_pointer(), None)
|
||
format_str: str = ""
|
||
cast_args: list[Any] = []
|
||
for part in Node.values:
|
||
if isinstance(part, ast.Constant) and isinstance(part.value, str):
|
||
format_str += part.value.replace('\\', '\\\\').replace('%', '%%').replace('\n', '\\n').replace('\t', '\\t').replace('\0', '\\0')
|
||
elif isinstance(part, ast.FormattedValue):
|
||
fmt_class: str | None = self.Trans.ExprHandler._get_var_class(part.value, Gen)
|
||
if fmt_class and Gen._has_function(f'{fmt_class}.__str__'):
|
||
obj_val: Any = self.HandleExprLlvm(part.value)
|
||
if obj_val:
|
||
if isinstance(obj_val.type, ir.PointerType) and isinstance(obj_val.type.pointee, ir.IntType) and obj_val.type.pointee.width == 8:
|
||
if fmt_class in Gen.structs:
|
||
obj_val = Gen.builder.bitcast(obj_val, ir.PointerType(Gen.structs[fmt_class]), name=f"cast_{fmt_class}")
|
||
val: Any = Gen.builder.call(Gen._get_function(f'{fmt_class}.__str__'), [obj_val], name=f"call_{fmt_class}.__str__")
|
||
Gen._RegisterTempPtr(val)
|
||
format_str += "%s"
|
||
cast_args.append(val)
|
||
else:
|
||
format_str += "%s"
|
||
cast_args.append(ir.Constant(ir.PointerType(ir.IntType(8)), None))
|
||
continue
|
||
val = self.HandleExprLlvm(part.value)
|
||
if not val:
|
||
format_str += "%d"
|
||
cast_args.append(ir.Constant(ir.IntType(32), 0))
|
||
continue
|
||
# 检查 format_spec(Python f-string 格式说明符)
|
||
if part.format_spec is not None:
|
||
c_fmt: str | None = self._parse_format_spec(part.format_spec, val)
|
||
if c_fmt is not None:
|
||
cast_val: Any = self._cast_value_for_format(val, c_fmt, Gen, part.value)
|
||
format_str += c_fmt
|
||
cast_args.append(cast_val)
|
||
continue
|
||
if isinstance(val.type, ir.PointerType):
|
||
pointee: ir.Type = val.type.pointee
|
||
if isinstance(pointee, ir.IntType) and pointee.width == 8:
|
||
format_str += "%s"
|
||
cast_args.append(val)
|
||
elif isinstance(pointee, (ir.LiteralStructType, ir.IdentifiedStructType)):
|
||
int_val: Any = Gen.builder.ptrtoint(val, ir.IntType(64), name="fstr_ptr2int")
|
||
trunc_val: Any = Gen.builder.trunc(int_val, ir.IntType(32), name="fstr_trunc")
|
||
format_str += "%d"
|
||
cast_args.append(trunc_val)
|
||
else:
|
||
Loaded: Any = Gen._load(val, name="fstr_deref")
|
||
if isinstance(Loaded.type, ir.IntType):
|
||
if Loaded.type.width == 8:
|
||
format_str += "%c"
|
||
ext: Any = Gen.builder.zext(Loaded, ir.IntType(32), name="fstr_char2int")
|
||
cast_args.append(ext)
|
||
else:
|
||
format_str += "%d"
|
||
cast_args.append(Loaded)
|
||
elif isinstance(Loaded.type, (ir.FloatType, ir.DoubleType)):
|
||
format_str += "%f"
|
||
cast_args.append(Loaded)
|
||
elif isinstance(Loaded.type, ir.PointerType):
|
||
format_str += "%s"
|
||
cast_args.append(Loaded)
|
||
else:
|
||
format_str += "%d"
|
||
cast_args.append(Loaded)
|
||
elif isinstance(val.type, ir.IntType):
|
||
if val.type.width == 8:
|
||
format_str += "%c"
|
||
ext = Gen.builder.zext(val, ir.IntType(32), name="fstr_char2int")
|
||
cast_args.append(ext)
|
||
elif val.type.width == 1:
|
||
zext: Any = Gen.builder.zext(val, ir.IntType(32), name="fstr_bool2int")
|
||
format_str += "%d"
|
||
cast_args.append(zext)
|
||
else:
|
||
is_u: bool = Gen._check_node_unsigned(part.value)
|
||
format_str += "%u" if is_u else "%d"
|
||
cast_args.append(val)
|
||
elif isinstance(val.type, (ir.FloatType, ir.DoubleType)):
|
||
format_str += "%f"
|
||
cast_args.append(val)
|
||
else:
|
||
format_str += "%d"
|
||
cast_args.append(val)
|
||
if not format_str:
|
||
return Gen.emit_constant("", 'string')
|
||
format_ptr: Any = Gen._create_string_global(format_str)
|
||
# 统一使用 snprintf 生成字符串,返回 i8*
|
||
# __str__ 方法内使用堆分配(返回值需持久化),其他场景使用栈分配(自动释放)
|
||
in_str_method: bool = bool(getattr(Gen, 'func', None) and getattr(Gen.func, 'name', '').endswith('.__str__'))
|
||
snprintf_type: ir.FunctionType = ir.FunctionType(ir.IntType(32), [
|
||
ir.PointerType(ir.IntType(8)),
|
||
ir.IntType(64),
|
||
ir.PointerType(ir.IntType(8))
|
||
], var_arg=True)
|
||
snprintf_func: Any = Gen.get_or_declare_c_func('snprintf', snprintf_type)
|
||
null_ptr: ir.Constant = ir.Constant(ir.PointerType(ir.IntType(8)), None)
|
||
# 第一遍:获取所需长度(NULL 缓冲,size=0)
|
||
len_val: ir.Value = Gen.builder.call(snprintf_func, [null_ptr, ir.Constant(ir.IntType(64), 0), format_ptr] + cast_args, name="fstr_len")
|
||
# 加 1 用于 null 终止符
|
||
size_val: ir.Value = Gen.builder.add(len_val, ir.Constant(ir.IntType(32), 1), name="fstr_size")
|
||
size_ext: ir.Value = Gen.builder.zext(size_val, ir.IntType(64), name="fstr_size_ext")
|
||
if in_str_method:
|
||
# 堆分配(__str__ 返回值需跨函数生命周期)
|
||
malloc_type: ir.FunctionType = ir.FunctionType(ir.PointerType(ir.IntType(8)), [ir.IntType(64)])
|
||
malloc_func: Any = Gen._get_or_declare_func('malloc', malloc_type)
|
||
buf_ptr: ir.Value = Gen.builder.call(malloc_func, [size_ext], name="fstr_heap_buf")
|
||
Gen.builder.call(snprintf_func, [buf_ptr, size_ext, format_ptr] + cast_args, name="fstr_snprintf")
|
||
Gen._RegisterTempPtr(buf_ptr)
|
||
return buf_ptr
|
||
# 栈分配(局部使用,自动释放)
|
||
buf_ptr: ir.Value = Gen._alloca(ir.IntType(8), name="fstr_stack_buf", size=size_val)
|
||
Gen.builder.call(snprintf_func, [buf_ptr, size_ext, format_ptr] + cast_args, name="fstr_snprintf")
|
||
return buf_ptr
|
||
|
||
def _parse_format_spec(self, format_spec_node: ast.AST, val: ir.Value) -> str | None:
|
||
"""将 Python f-string format_spec 转换为 C printf 格式说明符。
|
||
|
||
支持: d/x/X/o/c/f/s 及 [[fill]align][sign][#][0][width][.precision][type] 语法。
|
||
返回 None 表示无法解析,调用方应回退到类型推断。
|
||
"""
|
||
if isinstance(format_spec_node, ast.JoinedStr):
|
||
parts: list[str] = []
|
||
for v in format_spec_node.values:
|
||
if isinstance(v, ast.Constant) and isinstance(v.value, str):
|
||
parts.append(v.value)
|
||
else:
|
||
return None
|
||
spec_str: str = ''.join(parts)
|
||
elif isinstance(format_spec_node, ast.Constant) and isinstance(format_spec_node.value, str):
|
||
spec_str = format_spec_node.value
|
||
else:
|
||
return None
|
||
|
||
if not spec_str:
|
||
return None
|
||
|
||
m: Any = re.match(r'^(.?[<>=^])?([+\- ])?(#)?(0)?(\d+)?(,)?(\.\d+)?([bcdeEfFgGnosxX%])?$', spec_str)
|
||
if not m:
|
||
return None
|
||
|
||
fill_align, sign, alt, zero, width, comma, precision, fmt_type = m.groups()
|
||
|
||
c_flags: str = ''
|
||
if sign == '+':
|
||
c_flags += '+'
|
||
elif sign == ' ':
|
||
c_flags += ' '
|
||
if alt:
|
||
c_flags += '#'
|
||
|
||
if fill_align:
|
||
align_char: str = fill_align[-1] if len(fill_align) > 1 else fill_align[0]
|
||
if align_char in ('<', '^'):
|
||
c_flags += '-'
|
||
if zero:
|
||
c_flags += '0'
|
||
|
||
c_width: str = width or ''
|
||
c_precision: str = precision or ''
|
||
|
||
length_mod: str = ''
|
||
if isinstance(val.type, ir.IntType):
|
||
if val.type.width == 64:
|
||
length_mod = 'll'
|
||
elif val.type.width == 32:
|
||
length_mod = 'l'
|
||
elif val.type.width == 16:
|
||
length_mod = 'h'
|
||
elif val.type.width == 8:
|
||
length_mod = 'hh'
|
||
|
||
c_specifier: str = ''
|
||
if fmt_type in ('d', 'n'):
|
||
c_specifier = f'{length_mod}d'
|
||
elif fmt_type == 'x':
|
||
c_specifier = f'{length_mod}x'
|
||
elif fmt_type == 'X':
|
||
c_specifier = f'{length_mod}X'
|
||
elif fmt_type == 'o':
|
||
c_specifier = f'{length_mod}o'
|
||
elif fmt_type == 'b':
|
||
return None
|
||
elif fmt_type == 'c':
|
||
c_specifier = 'c'
|
||
elif fmt_type in ('f', 'F'):
|
||
c_specifier = 'f'
|
||
elif fmt_type in ('e', 'E', 'g', 'G'):
|
||
c_specifier = fmt_type
|
||
elif fmt_type == 's':
|
||
c_specifier = 's'
|
||
elif fmt_type == '%':
|
||
c_specifier = 'f%%'
|
||
else:
|
||
if isinstance(val.type, ir.IntType):
|
||
if val.type.width == 8:
|
||
c_specifier = 'c'
|
||
elif val.type.width == 1:
|
||
c_specifier = 'd'
|
||
else:
|
||
c_specifier = f'{length_mod}d'
|
||
elif isinstance(val.type, (ir.FloatType, ir.DoubleType)):
|
||
c_specifier = 'f'
|
||
elif isinstance(val.type, ir.PointerType):
|
||
c_specifier = 's'
|
||
else:
|
||
c_specifier = 'd'
|
||
|
||
return f'%{c_flags}{c_width}{c_precision}{c_specifier}'
|
||
|
||
def _cast_value_for_format(self, val: ir.Value, c_fmt: str, Gen: Any, arg_node: ast.AST) -> ir.Value:
|
||
"""根据 C printf 格式说明符转换值类型。"""
|
||
fmt_type_char: str = ''
|
||
for ch in reversed(c_fmt):
|
||
if ch.isalpha() and ch not in ('l', 'h'):
|
||
fmt_type_char = ch
|
||
break
|
||
elif ch == '%':
|
||
break
|
||
|
||
if fmt_type_char == 's':
|
||
if isinstance(val.type, ir.PointerType):
|
||
pointee: ir.Type = val.type.pointee
|
||
if isinstance(pointee, ir.IntType) and pointee.width == 8:
|
||
return val
|
||
if isinstance(pointee, ir.ArrayType) and isinstance(pointee.element, ir.IntType) and pointee.element.width == 8:
|
||
return Gen.builder.gep(val, [ir.Constant(ir.IntType(32), 0), ir.Constant(ir.IntType(32), 0)], name="fstr_arr2ptr")
|
||
return val
|
||
elif fmt_type_char == 'c':
|
||
if isinstance(val.type, ir.IntType) and val.type.width < 32:
|
||
return Gen.builder.zext(val, ir.IntType(32), name="fstr_zext_c")
|
||
return val
|
||
elif fmt_type_char in ('d', 'x', 'X', 'o', 'i'):
|
||
if isinstance(val.type, ir.IntType):
|
||
if val.type.width == 8:
|
||
is_u: bool = Gen._check_node_unsigned(arg_node)
|
||
if is_u:
|
||
return Gen.builder.zext(val, ir.IntType(32), name="fstr_zext_i8")
|
||
return Gen.builder.sext(val, ir.IntType(32), name="fstr_sext_i8")
|
||
elif val.type.width == 16:
|
||
is_u = Gen._check_node_unsigned(arg_node)
|
||
if is_u:
|
||
return Gen.builder.zext(val, ir.IntType(32), name="fstr_zext_i16")
|
||
return Gen.builder.sext(val, ir.IntType(32), name="fstr_sext_i16")
|
||
elif val.type.width == 1:
|
||
return Gen.builder.zext(val, ir.IntType(32), name="fstr_zext_bool")
|
||
return val
|
||
elif isinstance(val.type, (ir.FloatType, ir.DoubleType)):
|
||
return Gen.builder.fptosi(val, ir.IntType(32), name="fstr_f2i")
|
||
return val
|
||
elif fmt_type_char in ('f', 'e', 'E', 'g', 'G'):
|
||
if isinstance(val.type, ir.FloatType):
|
||
return Gen.builder.fpext(val, ir.DoubleType(), name="fstr_fpext")
|
||
elif isinstance(val.type, ir.IntType):
|
||
return Gen.builder.sitofp(val, ir.DoubleType(), name="fstr_i2f")
|
||
return val
|
||
return val |