pslang/libs/jit/source/arch/linux_x86_64/compiler.cpp

1342 lines
41 KiB
C++

#include <pslang/jit/arch/linux_x86_64/compiler.hpp>
#include <pslang/jit/arch/linux_x86_64/instruction_builder.hpp>
#include <pslang/jit/helpers.hpp>
#include <pslang/ir/node.hpp>
#include <pslang/ir/compiler.hpp>
#include <pslang/ast/type.hpp>
#include <pslang/ast/struct.hpp>
#include <pslang/ast/function.hpp>
#include <pslang/types/type_visitor.hpp>
namespace pslang::jit::linux_x86_64
{
namespace
{
template <typename ... Regs>
reg find_free_reg(reg start, Regs ... regs)
{
for (int i = (int)start; i < 16; ++i)
{
if ((true && ... && (regs != reg(i))))
return reg(i);
}
throw std::runtime_error("Unable to find a free register");
}
static std::array<reg, 6> integer_arg_reg = {
reg::rdi,
reg::rsi,
reg::rdx,
reg::rcx,
reg::r8,
reg::r9,
};
static std::array<reg, 2> integer_ret_reg = {
reg::rax,
reg::rdx,
};
struct small_struct_data
{
struct octet
{
bool has_integer = false;
bool has_floating_point = false;
};
std::optional<octet> octets[2];
};
struct value_address
{
// None means RIP-based addressing
std::optional<reg> base;
std::int32_t offset;
};
struct local_context
{
bool use_frame_pointer = true;
std::unordered_map<std::string, std::int32_t> extern_symbols;
std::unordered_map<types::type_ptr, small_struct_data> small_structs;
std::unordered_map<ir::node_ref, std::int32_t> nodes;
struct resolve_data
{
std::int32_t offset;
ir::node_ref target;
};
std::vector<resolve_data> jump_resolve;
std::vector<resolve_data> cjump_resolve;
std::vector<resolve_data> node_resolve;
std::vector<resolve_data> call_resolve;
};
std::optional<small_struct_data> classify_small_struct(local_context & lcontext, types::type_ptr type)
{
// TODO: unit type?
if (ast::type_size(*type) > 16)
return std::nullopt;
if (auto it = lcontext.small_structs.find(type); it != lcontext.small_structs.end())
return it->second;
if (std::holds_alternative<types::struct_type>(*type) || std::holds_alternative<types::array_type>(*type))
{
small_struct_data result;
auto layout = ast::flat_layout(*type);
for (auto const & field : layout)
{
std::size_t octet_index = field.offset / 8;
auto & octet = result.octets[octet_index].emplace();
if (types::is_floating_point_type(*field.type))
octet.has_floating_point = true;
else
octet.has_integer = true;
}
return (lcontext.small_structs[type] = result);
}
if (types::is_floating_point_type(*type))
return small_struct_data{.octets = {small_struct_data::octet{.has_floating_point = true}, std::nullopt}};
return small_struct_data{.octets = {small_struct_data::octet{.has_integer = true}, std::nullopt}};
}
struct populate_globals_visitor
{
std::vector<std::uint8_t> & storage;
local_context & lcontext;
template <typename Node>
void apply(ir::node_ref, Node const &, types::type_ptr const &)
{}
void apply(ir::node_ref it, ir::global const & node, types::type_ptr const & type)
{
auto size = ast::type_size(*type);
if (node.initializer.size() > size)
throw std::runtime_error("global IR node with initializer larger than type size");
auto alignment = ast::type_alignment(*type);
auto offset = storage.size();
offset = ((offset + (alignment - 1)) / alignment) * alignment;
storage.resize(offset + size);
std::copy(node.initializer.begin(), node.initializer.end(), storage.begin() + offset);
lcontext.nodes[it] = offset;
}
};
struct populate_const_data_visitor
{
program_context & pcontext;
local_context & lcontext;
template <typename Node>
void apply(Node const & node, types::type_ptr const &)
{}
void apply(ir::extern_symbol const & node, types::type_ptr const &)
{
std::int32_t offset = pcontext.storage.size();
lcontext.extern_symbols[node.name] = offset;
pcontext.foreign_resolve.push_back({node.name, offset});
push_bytes(pcontext.storage.storage, (void *)nullptr);
}
};
struct literal_visitor
{
program_context & pcontext;
local_context & lcontext;
instruction_builder & builder;
void operator()(ast::bool_literal const & node)
{
if (node.value)
builder.mov(-1, reg::rax);
else
builder.xor_(reg::rax, reg::rax);
}
template <typename T>
requires(std::is_integral_v<T> && !std::is_same_v<T, bool>)
void operator()(ast::primitive_literal_base<T> const & node)
{
if (sizeof(T) <= 4)
{
if (std::is_signed_v<T>)
builder.movsd(std::int32_t(node.value), reg::rax);
else
builder.movzd(std::uint32_t(node.value), reg::rax);
}
else
builder.mov(std::uint64_t(node.value), reg::rax);
}
void operator()(ast::f16_literal const & node)
{
throw std::runtime_error("Not implemented");
}
void operator()(ast::f32_literal const & node)
{
builder.movzd(*(std::uint32_t const *)(&node.value), reg::rax);
builder.mov_gpr_to_xmm_32(reg::rax, reg::xmm0);
}
void operator()(ast::f64_literal const & node)
{
builder.mov(*(std::uint64_t const *)(&node.value), reg::rax);
builder.mov_gpr_to_xmm(reg::rax, reg::xmm0);
}
};
struct compile_visitor
{
program_context & pcontext;
ir::module_context const & mcontext;
local_context & lcontext;
instruction_builder & builder;
std::vector<std::int32_t> argument_position;
std::unordered_map<ir::node_ref, std::int32_t> stack_position;
std::int32_t stack_size = 0;
bool return_value_is_large_struct = false;
void apply(ir::node_ref, ir::label const &, types::type_ptr const &)
{}
void apply(ir::node_ref it, ir::literal const & node, types::type_ptr const & type)
{
std::visit(literal_visitor{pcontext, lcontext, builder}, node.value);
if (types::is_integer_like_type(*type))
store(it, reg::rax);
else if (types::is_floating_point_type(*type))
store_xmm(it, reg::xmm0, ast::type_size(*type));
}
void apply(ir::node_ref it, ir::alloc const & node, types::type_ptr const & type)
{
// Nothing to do: alloc just allocates a node of some type,
// but we already allocated stack space for it
}
void apply(ir::node_ref it, ir::global const & node, types::type_ptr const & type)
{
// Globals are added in a separate pass before code
}
void apply(ir::node_ref it, ir::copy const & node, types::type_ptr const & type)
{
auto size = ast::type_size(*type);
auto dst_address = node_address(it);
auto src_type = node.source->inferred_type;
auto src_address = node_address(node.source);
for (auto field_id : node.path)
{
if (auto struct_type = std::get_if<types::struct_type>(src_type.get()))
{
auto struct_node = struct_type->node;
src_type = struct_node->fields[field_id].inferred_type;
src_address.offset += struct_node->fields[field_id].layout.offset;
}
else if (auto array_type = std::get_if<types::array_type>(src_type.get()))
{
src_type = array_type->element_type;
src_address.offset += field_id * ast::type_size(*array_type->element_type);
}
else
throw std::runtime_error("Unknown object type for field copy");
}
copy_memory(src_address, dst_address, size);
}
void apply(ir::node_ref it, ir::load const & node, types::type_ptr const & type)
{
load(node.ptr, reg::rax);
auto size = ast::type_size(*type);
auto dst_address = node_address(it);
copy_memory({.base = reg::rax, .offset = 0}, dst_address, size);
}
void apply(ir::node_ref it, ir::store const & node, types::type_ptr const & type)
{
load(node.ptr, reg::rax);
auto size = ast::type_size(*type);
auto src_address = node_address(node.value);
copy_memory(src_address, {.base = reg::rax, .offset = 0}, size);
}
void apply(ir::node_ref it, ir::unary_operation const & node, types::type_ptr const & type)
{
switch (node.type)
{
case ast::unary_operation_type::negation:
if (types::is_integer_type(*type))
{
load(node.arg1, reg::rax);
builder.neg(reg::rax);
store(it, reg::rax);
}
else if (types::is_floating_point_type(*type))
{
auto const size = ast::type_size(*type);
load_xmm(node.arg1, reg::xmm0, size);
if (size <= 4)
{
// f16 or f32
builder.movzd(0x80000000u, reg::rax);
builder.mov_gpr_to_xmm_32(reg::rax, reg::xmm1);
builder.xor_xmm_32(reg::xmm1, reg::xmm0);
}
else
{
// f64
builder.mov(0x8000000000000000ull, reg::rax);
builder.mov_gpr_to_xmm(reg::rax, reg::xmm1);
builder.xor_xmm(reg::xmm1, reg::xmm0);
}
store_xmm(it, reg::xmm0, size);
}
break;
case ast::unary_operation_type::logical_not:
load(node.arg1, reg::rax);
builder.not_(reg::rax);
store(it, reg::rax);
break;
case ast::unary_operation_type::address_of:
case ast::unary_operation_type::mutable_address_of:
{
auto const address = node_address(node.arg1);
if (address.base)
builder.lea(*address.base, address.offset, reg::rax);
else
builder.lea_rip_prev(address.offset, reg::rax);
store(it, reg::rax);
}
break;
case ast::unary_operation_type::dereference:
throw std::runtime_error("Dereference operator mush not be present in compiled IR");
}
}
void apply(ir::node_ref it, ir::binary_operation const & node, types::type_ptr const & type)
{
auto const arg1_type = node.arg1->inferred_type;
auto const arg2_type = node.arg2->inferred_type;
bool const is_fp = types::is_floating_point_type(*arg1_type);
bool const result_is_fp = types::is_floating_point_type(*type);
auto const arg_size = ast::type_size(*arg1_type);
if (is_fp)
{
load_xmm(node.arg1, reg::xmm0, arg_size);
load_xmm(node.arg2, reg::xmm1, arg_size);
}
else
{
load(node.arg1, reg::rax);
load(node.arg2, reg::rbx);
}
switch (node.type)
{
case ast::binary_operation_type::addition:
if (is_fp)
{
if (arg_size <= 4)
builder.add_xmm_32(reg::xmm1, reg::xmm0);
else
builder.add_xmm(reg::xmm1, reg::xmm0);
}
else
builder.add(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::subtraction:
if (is_fp)
{
if (arg_size <= 4)
builder.sub_xmm_32(reg::xmm1, reg::xmm0);
else
builder.sub_xmm(reg::xmm1, reg::xmm0);
}
else
builder.sub(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::multiplication:
if (is_fp)
{
if (arg_size <= 4)
builder.mul_xmm_32(reg::xmm1, reg::xmm0);
else
builder.mul_xmm(reg::xmm1, reg::xmm0);
}
else
builder.mul(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::division:
if (is_fp)
{
if (arg_size <= 4)
builder.div_xmm_32(reg::xmm1, reg::xmm0);
else
builder.div_xmm(reg::xmm1, reg::xmm0);
}
else
{
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
if (types::is_unsigned_integer_type(*arg1_type))
{
builder.xor_(reg::rdx, reg::rdx);
builder.udiv(reg::rbx);
}
else
{
builder.cqo();
builder.idiv(reg::rbx);
}
}
break;
case ast::binary_operation_type::remainder:
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
if (types::is_unsigned_integer_type(*arg1_type))
{
builder.xor_(reg::rdx, reg::rdx);
builder.udiv(reg::rbx);
builder.mov(reg::rdx, reg::rax);
}
else
{
builder.cqo();
builder.idiv(reg::rbx);
builder.mov(reg::rdx, reg::rax);
}
break;
case ast::binary_operation_type::binary_and:
builder.and_(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::logical_and:
throw std::runtime_error("Short-circuiting operators must have been unwrapped in IR compiler");
case ast::binary_operation_type::binary_or:
builder.or_(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::logical_or:
throw std::runtime_error("Short-circuiting operators must have been unwrapped in IR compiler");
case ast::binary_operation_type::logical_xor:
builder.xor_(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::left_shift:
builder.mov(reg::rbx, reg::rcx);
builder.shl(reg::rax);
break;
case ast::binary_operation_type::right_shift:
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rcx, *arg2_type);
if (types::is_unsigned_integer_type(*arg1_type))
builder.shr(reg::rax);
else
builder.sar(reg::rax);
break;
case ast::binary_operation_type::equals:
if (is_fp)
{
if (arg_size <= 4)
builder.cmp_xmm_32(reg::xmm0, reg::xmm1);
else
builder.cmp_xmm(reg::xmm0, reg::xmm1);
}
else
{
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
builder.cmp(reg::rax, reg::rbx);
}
builder.mov(0, reg::rax);
builder.mov(-1, reg::rbx);
builder.cmovz(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::not_equals:
if (is_fp)
{
if (arg_size <= 4)
builder.cmp_xmm_32(reg::xmm0, reg::xmm1);
else
builder.cmp_xmm(reg::xmm0, reg::xmm1);
}
else
{
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
builder.cmp(reg::rax, reg::rbx);
}
builder.mov(0, reg::rax);
builder.mov(-1, reg::rbx);
builder.cmovnz(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::less:
if (is_fp)
{
if (arg_size <= 4)
builder.cmp_xmm_32(reg::xmm0, reg::xmm1);
else
builder.cmp_xmm(reg::xmm0, reg::xmm1);
}
else
{
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
builder.cmp(reg::rax, reg::rbx);
}
builder.mov(0, reg::rax);
builder.mov(-1, reg::rbx);
if (is_fp || types::is_bool_type(*arg1_type) || types::is_unsigned_integer_type(*arg1_type))
builder.cmovb(reg::rbx, reg::rax);
else
builder.cmovl(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::greater:
if (is_fp)
{
if (arg_size <= 4)
builder.cmp_xmm_32(reg::xmm1, reg::xmm0);
else
builder.cmp_xmm(reg::xmm1, reg::xmm0);
}
else
{
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
builder.cmp(reg::rbx, reg::rax);
}
builder.mov(0, reg::rax);
builder.mov(-1, reg::rbx);
if (is_fp || types::is_bool_type(*arg1_type) || types::is_unsigned_integer_type(*arg1_type))
builder.cmovb(reg::rbx, reg::rax);
else
builder.cmovl(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::less_equals:
if (is_fp)
{
if (arg_size <= 4)
builder.cmp_xmm_32(reg::xmm1, reg::xmm0);
else
builder.cmp_xmm(reg::xmm1, reg::xmm0);
}
else
{
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
builder.cmp(reg::rbx, reg::rax);
}
builder.mov(0, reg::rax);
builder.mov(-1, reg::rbx);
if (is_fp || types::is_bool_type(*arg1_type) || types::is_unsigned_integer_type(*arg1_type))
builder.cmovnb(reg::rbx, reg::rax);
else
builder.cmovnl(reg::rbx, reg::rax);
break;
case ast::binary_operation_type::greater_equals:
if (is_fp)
{
if (arg_size <= 4)
builder.cmp_xmm_32(reg::xmm0, reg::xmm1);
else
builder.cmp_xmm(reg::xmm0, reg::xmm1);
}
else
{
reg_extend(reg::rax, reg::rax, *arg1_type);
reg_extend(reg::rbx, reg::rbx, *arg1_type);
builder.cmp(reg::rax, reg::rbx);
}
builder.mov(0, reg::rax);
builder.mov(-1, reg::rbx);
if (is_fp || types::is_bool_type(*arg1_type) || types::is_unsigned_integer_type(*arg1_type))
builder.cmovnb(reg::rbx, reg::rax);
else
builder.cmovnl(reg::rbx, reg::rax);
break;
}
if (result_is_fp)
store_xmm(it, reg::xmm0, arg_size);
else
store(it, reg::rax);
}
void apply(ir::node_ref it, ir::cast_operation const & node, types::type_ptr const & type)
{
auto src_type = node.arg1->inferred_type;
auto dst_type = node.target_type;
auto u64_type = types::primitive_type{types::u64_type{}};
if (false
|| (types::is_builtin_type(*src_type) && types::equal(*src_type, *dst_type))
|| (types::is_pointer_type(*src_type) && types::is_pointer_type(*dst_type))
|| (types::equal(*src_type, u64_type) && types::is_pointer_type(*dst_type))
|| (types::equal(*dst_type, u64_type) && types::is_pointer_type(*src_type))
)
{
load(node.arg1, reg::rax);
store(it, reg::rax);
return;
}
if (auto array_type = std::get_if<types::array_type>(src_type.get()))
{
if (auto pointer_type = std::get_if<types::pointer_type>(dst_type.get()))
{
if (types::equal(*array_type->element_type, *pointer_type->referenced_type))
{
auto arg1_address = node_address(node.arg1);
if (arg1_address.base)
builder.lea(*arg1_address.base, arg1_address.offset, reg::rax);
else
builder.lea_rip_prev(arg1_address.offset, reg::rax);
store(it, reg::rax);
return;
}
}
}
if (types::is_numeric_type(*src_type) && types::is_numeric_type(*dst_type))
{
if (types::is_integer_type(*src_type))
{
load(node.arg1, reg::rax);
if (types::is_integer_type(*dst_type))
{
reg_extend(reg::rax, reg::rax, *src_type);
}
else if (types::is_floating_point_type(*dst_type))
{
reg_extend(reg::rax, reg::rax, *src_type);
auto dst_size = ast::type_size(*dst_type);
if (types::is_unsigned_integer_type(*src_type) && ast::type_size(*src_type) == 8)
{
// No u64 -> f32/f64 conversion on x86_64
// Emulate it by splitting x = 2 * (x >> 1) + (x & 1)
reg_extend(reg::rax, reg::rax, *src_type);
builder.test(reg::rax, reg::rax);
auto first_jump = builder.code.size();
builder.jump_if_sign(0);
if (dst_size <= 4)
builder.gpr_to_xmm_32(reg::rax, reg::xmm0);
else
builder.gpr_to_xmm(reg::rax, reg::xmm0);
auto second_jump = builder.code.size();
builder.jump(0);
auto first_jump_target = builder.code.size();
builder.mov(reg::rax, reg::rcx);
builder.shr1(reg::rcx);
builder.and_imm8(reg::rax, 1);
builder.or_(reg::rax, reg::rcx);
if (dst_size <= 4)
{
builder.gpr_to_xmm_32(reg::rcx, reg::xmm0);
builder.add_xmm_32(reg::xmm0, reg::xmm0);
}
else
{
builder.gpr_to_xmm(reg::rcx, reg::xmm0);
builder.add_xmm(reg::xmm0, reg::xmm0);
}
auto second_jump_target = builder.code.size();
builder.cjump_inject_prev(builder.code.data() + first_jump, first_jump_target - first_jump);
builder.jump_inject_prev(builder.code.data() + second_jump, second_jump_target - second_jump);
}
else
{
if (dst_size <= 4)
builder.gpr_to_xmm_32(reg::rax, reg::xmm0);
else
builder.gpr_to_xmm(reg::rax, reg::xmm0);
}
}
}
else if (types::is_floating_point_type(*src_type))
{
if (types::is_integer_type(*dst_type))
{
// TODO: correct floating-point to 64-bit unsigned conversion
auto src_size = ast::type_size(*src_type);
if (src_size <= 4)
builder.xmm_32_to_gpr(reg::xmm0, reg::rax);
else
builder.xmm_to_gpr(reg::xmm0, reg::rax);
}
else if (types::is_floating_point_type(*dst_type))
{
auto src_size = ast::type_size(*src_type);
auto dst_size = ast::type_size(*dst_type);
if (src_size <= 4 && dst_size == 8)
builder.xmm_32_to_64(reg::xmm0, reg::xmm0);
else if (src_size == 8 && dst_size <= 4)
builder.xmm_64_to_32(reg::xmm0, reg::xmm0);
}
}
if (types::is_integer_type(*dst_type))
{
store(it, reg::rax);
}
else if (types::is_floating_point_type(*dst_type))
{
store_xmm(it, reg::xmm0, ast::type_size(*dst_type));
}
return;
}
throw std::runtime_error("Unknown types for cast instruction");
}
void apply(ir::node_ref, ir::argument const &, types::type_ptr const &)
{
// Nothing to do: arguments already pushed on stack in function preamble
}
void apply(ir::node_ref it, ir::instruction_address const & node, types::type_ptr const & type)
{
lcontext.node_resolve.emplace_back(pcontext.storage.size(), node.target);
builder.lea_rip(0, reg::rax);
store(it, reg::rax);
}
void apply(ir::node_ref it, ir::extern_symbol const & node, types::type_ptr const & type)
{
builder.mov_read_rip_prev(lcontext.extern_symbols[node.name] - (std::int32_t)pcontext.storage.size(), reg::rax);
store(it, reg::rax);
}
void apply(ir::node_ref, ir::assignment const & node, types::type_ptr const & type)
{
auto src_address = node_address(node.rhs);
auto dst_type = node.lhs->inferred_type;
auto dst_address = node_address(node.lhs);
for (auto field_id : node.path)
{
if (auto struct_type = std::get_if<types::struct_type>(dst_type.get()))
{
auto struct_node = struct_type->node;
dst_type = struct_node->fields[field_id].inferred_type;
dst_address.offset += struct_node->fields[field_id].layout.offset;
}
else if (auto array_type = std::get_if<types::array_type>(dst_type.get()))
{
dst_type = array_type->element_type;
dst_address.offset += field_id * ast::type_size(*array_type->element_type);
}
else
throw std::runtime_error("Unknown object type for field assignment");
}
copy_memory(src_address, dst_address, ast::type_size(*dst_type));
}
void apply(ir::node_ref, ir::jump const & node, types::type_ptr const & type)
{
lcontext.jump_resolve.push_back({(std::int32_t)pcontext.storage.size(), node.target});
builder.jump(0);
}
void apply(ir::node_ref, ir::jump_if_zero const & node, types::type_ptr const & type)
{
load(node.condition, reg::rax);
reg_extend(reg::rax, reg::rax, *node.condition->inferred_type);
builder.test(reg::rax, reg::rax);
lcontext.cjump_resolve.push_back({(std::int32_t)pcontext.storage.size(), node.target});
builder.jump_if_zero(0);
}
void apply(ir::node_ref, ir::jump_if_nonzero const & node, types::type_ptr const & type)
{
load(node.condition, reg::rax);
reg_extend(reg::rax, reg::rax, *node.condition->inferred_type);
builder.test(reg::rax, reg::rax);
lcontext.cjump_resolve.push_back({(std::int32_t)pcontext.storage.size(), node.target});
builder.jump_if_nonzero(0);
}
template <typename Node, typename DoCall>
void apply_call(ir::node_ref it, Node const & node, types::type_ptr const & type, DoCall && do_call)
{
bool return_value_is_large_struct = false;
if (std::holds_alternative<types::struct_type>(*type) || std::holds_alternative<types::array_type>(*type))
return_value_is_large_struct = !classify_small_struct(lcontext, type);
if (this->return_value_is_large_struct)
{
builder.push(reg::rdi);
stack_size += 8;
}
// Compute space required for stack arguments
std::int32_t arguments_stack_size = 0;
for (auto const & argument : node.arguments)
{
auto size = ast::type_size(*argument->inferred_type);
auto alignment = ast::type_size(*argument->inferred_type);
auto struct_type = std::get_if<types::struct_type>(argument->inferred_type.get());
auto array_type = std::get_if<types::array_type>(argument->inferred_type.get());
if (struct_type || array_type)
{
if (!classify_small_struct(lcontext, argument->inferred_type))
{
arguments_stack_size += (alignment - (arguments_stack_size % alignment)) % alignment;
arguments_stack_size += size;
}
}
}
// Ensure stack is 16-byte aligned after CALL pushes return address on stack
// I.e. enforce (stack_size + arguments_stack_size) % 16 == 8
arguments_stack_size += (24 - ((arguments_stack_size + stack_size) % 16)) % 16;
// TODO: handle the case when there weren't enough registers
std::uint8_t reg_index = return_value_is_large_struct ? 1 : 0;
std::uint8_t fp_reg = 0;
std::int32_t stack_offset = 0;
for (auto const & argument : node.arguments)
{
auto size = ast::type_size(*argument->inferred_type);
auto alignment = ast::type_alignment(*argument->inferred_type);
auto struct_type = std::get_if<types::struct_type>(argument->inferred_type.get());
auto array_type = std::get_if<types::array_type>(argument->inferred_type.get());
if (struct_type || array_type)
{
if (auto small_struct = classify_small_struct(lcontext, argument->inferred_type))
{
auto address = node_address(argument);
for (int o : {0, 1})
{
auto & octet = small_struct->octets[o];
if (!octet) break;
if (octet->has_integer)
load(argument, integer_arg_reg[reg_index++], 8 * o);
else if (octet->has_floating_point)
load_xmm(argument, (reg)(fp_reg++), 8, 8 * o);
}
}
else
{
stack_offset += (alignment - (stack_offset % alignment)) % alignment;
copy_memory(node_address(argument), {.base = reg::rsp, .offset = stack_offset - arguments_stack_size}, size);
stack_offset += size;
}
}
else if (types::is_integer_like_type(*argument->inferred_type))
load(argument, integer_arg_reg[reg_index++]);
else if (types::is_floating_point_type(*argument->inferred_type))
load_xmm(argument, (reg)(fp_reg++), size);
else
throw std::runtime_error("Unsupported function argument type");
}
if (return_value_is_large_struct)
{
auto address = node_address(it);
// Function call node cannot have RIP-relative address
builder.lea(address.base.value(), address.offset, reg::rdi);
}
if (arguments_stack_size > 0)
{
builder.sub_imm(reg::rsp, arguments_stack_size);
stack_size += arguments_stack_size;
}
do_call();
if (arguments_stack_size > 0)
{
builder.add_imm(reg::rsp, arguments_stack_size);
stack_size -= arguments_stack_size;
}
std::size_t const size = ast::type_size(*type);
auto struct_type = std::get_if<types::struct_type>(type.get());
auto array_type = std::get_if<types::array_type>(type.get());
if (size == 0)
{}
else if (struct_type || array_type)
{
if (auto small_struct = classify_small_struct(lcontext, type))
{
std::uint8_t reg_index = 0;
std::uint8_t fp_reg = 0;
for (int o : {0, 1})
{
auto & octet = small_struct->octets[o];
if (!octet) break;
if (octet->has_integer)
store(it, integer_ret_reg[reg_index++], 8 * o);
else if (octet->has_floating_point)
store_xmm(it, (reg)(fp_reg++), 8, 8 * o);
}
}
else
{
// Nothing to be done - return value should already be in place
}
}
else if (types::is_integer_like_type(*type))
store(it, reg::rax);
else if (types::is_floating_point_type(*type))
store_xmm(it, reg::xmm0, size);
else
throw std::runtime_error("Unsupported return value type");
if (this->return_value_is_large_struct)
{
builder.pop(reg::rdi);
stack_size -= 8;
}
}
void apply(ir::node_ref it, ir::call const & node, types::type_ptr const & type)
{
apply_call(it, node, type, [&]{
lcontext.call_resolve.emplace_back(pcontext.storage.size(), node.target);
builder.call_imm(0);
});
}
void apply(ir::node_ref it, ir::call_pointer const & node, types::type_ptr const & type)
{
apply_call(it, node, type, [&]{
load(node.pointer, reg::rax);
builder.call_reg(reg::rax);
});
}
void apply(ir::node_ref, ir::return_value const & node, types::type_ptr const & type)
{
if (node.value)
{
auto type = (*node.value)->inferred_type;
auto size = ast::type_size(*type);
auto struct_type = std::get_if<types::struct_type>(type.get());
auto array_type = std::get_if<types::array_type>(type.get());
if (size == 0)
{}
else if (struct_type || array_type)
{
if (auto small_struct = classify_small_struct(lcontext, type))
{
auto address = node_address(*node.value);
std::uint8_t reg_index = 0;
std::uint8_t fp_reg = 0;
for (int o : {0, 1})
{
auto & octet = small_struct->octets[o];
if (!octet) break;
if (octet->has_integer)
load(*node.value, integer_ret_reg[reg_index++], 8 * o);
else if (octet->has_floating_point)
load_xmm(*node.value, (reg)(fp_reg++), 8, 8 * o);
}
}
else
{
copy_memory(node_address(*node.value), {reg::rdi, 0}, size);
}
}
else if (types::is_integer_like_type(*type))
load(*node.value, reg::rax);
else if (types::is_floating_point_type(*type))
load_xmm(*node.value, reg::xmm0, ast::type_size(*type));
else
throw std::runtime_error("Unsupported return value type");
}
if (stack_size > 0)
builder.add_imm(reg::rsp, stack_size);
builder.pop(reg::rbx);
if (lcontext.use_frame_pointer)
builder.pop(reg::rbp);
builder.ret();
}
void compile(ast::function_definition const * function_definition, ir::node_ref begin, ir::node_ref end)
{
auto result_type = function_definition->inferred_result_type;
if (result_type)
if (std::holds_alternative<types::struct_type>(*result_type) || std::holds_alternative<types::array_type>(*result_type))
return_value_is_large_struct = !classify_small_struct(lcontext, result_type);
stack_size = 0;
for (auto const & argument : function_definition->arguments)
{
auto size = ast::type_size(*argument.inferred_type);
// Ensure max alignment for simplicity
stack_size += ((size + 7) / 8) * 8;
argument_position.push_back(stack_size);
}
for (auto it = begin; it != end; ++it)
{
if (auto argument = std::get_if<ir::argument>(&it->instruction))
{
stack_position[it] = argument_position[argument->index];
}
else if (std::holds_alternative<ir::global>(it->instruction))
{
// stack position doesn't make sense for globals
}
else if (ir::is_value_instruction(it->instruction))
{
auto size = ast::type_size(*it->inferred_type);
if (size > 0)
stack_size += ((size + 7) / 8) * 8;
stack_position[it] = stack_size;
}
}
if (!std::holds_alternative<ir::label>(begin->instruction))
throw std::runtime_error("First IR node of a function must be a label");
auto it = begin;
lcontext.nodes[it] = pcontext.storage.size();
if (lcontext.use_frame_pointer)
{
builder.push(reg::rbp);
builder.mov(reg::rsp, reg::rbp);
}
builder.push(reg::rbx);
if (stack_size > 0)
builder.sub_imm(reg::rsp, stack_size);
// TODO: handle the case when there weren't enough registers
// If return value is large struct, RDI stores the pointer to return value
std::uint8_t reg_index = return_value_is_large_struct ? 1 : 0;
std::uint8_t fp_reg = 0;
// 8 for the return address
// Another 8 for RBP if using frame pointers
// Another 8 for saved RBX
std::int32_t stack_offset = stack_size + (lcontext.use_frame_pointer ? 16 : 8) + 8;
for (std::size_t i = 0; i < function_definition->arguments.size(); ++i)
{
auto const & argument = function_definition->arguments[i];
auto size = ast::type_size(*argument.inferred_type);
auto alignment = ast::type_alignment(*argument.inferred_type);
auto struct_type = std::get_if<types::struct_type>(argument.inferred_type.get());
auto array_type = std::get_if<types::array_type>(argument.inferred_type.get());
if (size == 0) continue;
if (struct_type || array_type)
{
if (auto small_struct = classify_small_struct(lcontext, argument.inferred_type))
{
for (std::size_t o : {0, 1})
{
auto & octet = small_struct->octets[o];
if (!octet) break;
if (octet->has_integer)
builder.mov_write(integer_arg_reg[reg_index++], reg::rsp, stack_size - argument_position[i] + 8 * o);
else if (octet->has_floating_point)
builder.mov_write_xmm((reg)(fp_reg++), reg::rsp, stack_size - argument_position[i] + 8 * o);
}
}
else
{
stack_offset += (alignment - (stack_offset % alignment)) % alignment;
copy_memory({.base = reg::rsp, .offset = stack_offset}, {.base = reg::rsp, .offset = stack_size - argument_position[i]}, size);
stack_offset += size;
}
}
else if (types::is_integer_like_type(*argument.inferred_type))
{
builder.mov_write(integer_arg_reg[reg_index++], reg::rsp, stack_size - argument_position[i]);
}
else if (types::is_floating_point_type(*argument.inferred_type))
{
auto const size = ast::type_size(*argument.inferred_type);
if (size == 4)
builder.mov_write_xmm_32((reg)(fp_reg++), reg::rsp, stack_size - argument_position[i]);
else if (size == 8)
builder.mov_write_xmm((reg)(fp_reg++), reg::rsp, stack_size - argument_position[i]);
else
throw std::runtime_error("Bad type size for floating-point argument");
}
else
throw std::runtime_error("Unknown argument type");
}
++it;
for (; it != end; ++it)
{
// Globals are added in a separate pass before code
if (lcontext.nodes.contains(it))
continue;
// Uncomment to debug per-node instruction generation:
builder.nop();
lcontext.nodes[it] = pcontext.storage.size();
std::visit([&](auto const & instruction){ apply(it, instruction, it->inferred_type); }, it->instruction);
}
}
private:
void reg_extend(reg reg_src, reg reg_dst, types::type const & type)
{
auto const size = ast::type_size(type);
if (types::is_unsigned_integer_type(type))
{
if (size == 1)
builder.movzx_8(reg_src, reg_dst);
else if (size == 2)
builder.movzx_16(reg_src, reg_dst);
else if (size == 4)
builder.movzx_32(reg_src, reg_dst);
else if (reg_src != reg_dst)
builder.mov(reg_src, reg_dst);
}
else if (types::is_signed_integer_type(type))
{
if (size == 1)
builder.movsx_8(reg_src, reg_dst);
else if (size == 2)
builder.movsx_16(reg_src, reg_dst);
else if (size == 4)
builder.movsx_32(reg_src, reg_dst);
else if (reg_src != reg_dst)
builder.mov(reg_src, reg_dst);
}
else if (types::is_bool_type(type))
{
builder.movzx_8(reg_src, reg_dst);
}
}
value_address node_address(ir::node_ref it)
{
if (std::holds_alternative<ir::global>(it->instruction))
return {.base = std::nullopt, .offset = lcontext.nodes.at(it) - static_cast<std::int32_t>(builder.code.size())};
else
return {.base = reg::rsp, .offset = stack_size - stack_position.at(it)};
}
void load(ir::node_ref it, reg reg_dst, std::int32_t offset = 0)
{
auto const address = node_address(it);
if (address.base)
builder.mov_read(*address.base, address.offset + offset, reg_dst);
else
builder.mov_read_rip_prev(address.offset + offset, reg_dst);
}
void store(ir::node_ref it, reg reg_src, std::int32_t offset = 0)
{
auto const address = node_address(it);
if (address.base)
builder.mov_write(reg_src, *address.base, address.offset + offset);
else
builder.mov_write_rip_prev(reg_src, address.offset + offset);
}
void load_xmm(ir::node_ref it, reg reg_dst, std::uint8_t size, std::int32_t offset = 0)
{
if (size != 4 && size != 8)
throw std::runtime_error("Bad type size for load_xmm");
auto const address = node_address(it);
if (address.base)
{
if (size == 4)
builder.mov_read_xmm_32(*address.base, address.offset + offset, reg_dst);
else
builder.mov_read_xmm(*address.base, address.offset + offset, reg_dst);
}
else
throw std::runtime_error("RIP-relative XMM read is not supported");
}
void store_xmm(ir::node_ref it, reg reg_src, std::uint8_t size, std::int32_t offset = 0)
{
if (size != 4 && size != 8)
throw std::runtime_error("Bad type size for store_xmm");
auto const address = node_address(it);
if (address.base)
{
if (size == 4)
builder.mov_write_xmm_32(reg_src, *address.base, address.offset + offset);
else
builder.mov_write_xmm(reg_src, *address.base, address.offset + offset);
}
else
throw std::runtime_error("RIP-relative XMM write is not supported");
}
void copy_memory(value_address src, value_address dst, std::size_t size)
{
auto const storage_size_at_start = (std::int32_t)pcontext.storage.size();
reg reg_src;
if (src.base)
reg_src = *src.base;
else
{
reg_src = find_free_reg(reg::r10, dst.base);
builder.lea_rip_prev(0, reg_src);
}
reg reg_dst;
if (dst.base)
reg_dst = *dst.base;
else
{
reg_dst = find_free_reg(reg::r10, reg_src);
builder.lea_rip_prev(storage_size_at_start - (std::int32_t)pcontext.storage.size(), reg_dst);
}
std::int32_t offset = 0;
auto copy_reg = find_free_reg(reg::r10, reg_src, reg_dst);
while (size > 0)
{
auto check_step = [&](std::size_t step)
{
return size >= step;
};
if (check_step(8))
{
builder.mov_read(reg_src, src.offset + offset, copy_reg);
builder.mov_write(copy_reg, reg_dst, dst.offset + offset);
size -= 8;
offset += 8;
}
else if (check_step(4))
{
builder.mov_read_32(reg_src, src.offset + offset, copy_reg);
builder.mov_write_32(copy_reg, reg_dst, dst.offset + offset);
size -= 4;
offset += 4;
}
else if (check_step(2))
{
builder.mov_read_16(reg_src, src.offset + offset, copy_reg);
builder.mov_write_16(copy_reg, reg_dst, dst.offset + offset);
size -= 2;
offset += 2;
}
else
{
builder.mov_read_8(reg_src, src.offset + offset, copy_reg);
builder.mov_write_8(copy_reg, reg_dst, dst.offset + offset);
size -= 1;
offset += 1;
}
}
}
};
}
void compile(program_context & pcontext, ir::module_context const & mcontext)
{
local_context lcontext;
auto data_begin = pcontext.storage.align();
{
populate_globals_visitor visitor{.storage = pcontext.storage.storage, .lcontext = lcontext};
for (auto it = mcontext.nodes->begin(); it != mcontext.nodes->end(); ++it)
std::visit([&](auto const & instruction){ visitor.apply(it, instruction, it->inferred_type); }, it->instruction);
}
#ifndef NDEBUG
// Force the data page to be allocated
// Helps with debugging in QtCreator (which pokes ~16 bytes before the code
// page and can hit unmapped memory)
if (data_begin == pcontext.storage.size())
pcontext.storage.storage.push_back(0);
#endif
auto data_end = pcontext.storage.align();
if (data_begin != data_end)
pcontext.storage.data.push_back({data_begin, data_end});
auto code_begin = data_end;
{
populate_const_data_visitor visitor{pcontext, lcontext};
for (auto it = mcontext.nodes->begin(); it != mcontext.nodes->end(); ++it)
std::visit([&](auto const & instruction){ visitor.apply(instruction, it->inferred_type); }, it->instruction);
}
instruction_builder builder{pcontext.storage.storage};
for (auto const & symbol : mcontext.symbols)
{
pcontext.symbols[symbol.first] = pcontext.storage.size();
compile_visitor visitor{pcontext, mcontext, lcontext, builder};
visitor.compile(symbol.first, symbol.second.begin, symbol.second.end);
}
pcontext.entry_point = lcontext.nodes.at(mcontext.entry_point);
for (auto const & resolve : lcontext.jump_resolve)
builder.jump_inject_prev(pcontext.storage.storage.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset);
for (auto const & resolve : lcontext.cjump_resolve)
builder.cjump_inject_prev(pcontext.storage.storage.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset);
for (auto const & resolve : lcontext.node_resolve)
builder.lea_rip_inject_prev(pcontext.storage.storage.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset);
for (auto const & resolve : lcontext.call_resolve)
builder.call_imm_inject_prev(pcontext.storage.storage.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset);
auto code_end = pcontext.storage.align();
pcontext.storage.code.push_back({code_begin, code_end});
}
}