#include #include #include #include #include #include #include #include #include namespace pslang::jit::linux_x86_64 { namespace { template 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 integer_arg_reg = { reg::rdi, reg::rsi, reg::rdx, reg::rcx, reg::r8, reg::r9, }; static std::array integer_ret_reg = { reg::rax, reg::rdx, }; struct small_struct_data { struct octet { bool has_integer = false; bool has_floating_point = false; }; std::optional octets[2]; }; struct value_address { std::optional base = std::nullopt; ir::node_ref global = {}; std::int32_t offset; }; bool is_global(ir::node_ref it) { return std::holds_alternative(it->instruction); } struct local_context { bool use_frame_pointer = true; std::unordered_map small_structs; std::unordered_map nodes; struct resolve_data { std::int32_t offset; ir::node_ref target; }; std::vector jump_resolve; std::vector cjump_resolve; std::vector node_resolve; std::vector call_resolve; }; std::optional 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(*type) || std::holds_alternative(*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 { compiled_module & result; local_context & lcontext; template 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 = result.data.size(); offset = ((offset + (alignment - 1)) / alignment) * alignment; result.data.resize(offset + size); std::copy(node.initializer.begin(), node.initializer.end(), result.data.begin() + offset); result.globals[it] = offset; } }; struct literal_visitor { 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 requires(std::is_integral_v && !std::is_same_v) void operator()(ast::primitive_literal_base const & node) { if (sizeof(T) <= 4) { if (std::is_signed_v) 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 { compiled_module & module; local_context & lcontext; instruction_builder & builder; std::vector argument_position; std::unordered_map 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{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(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(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 { push_resolve_global(address.global); builder.lea_rip_prev(0, 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(src_type.get())) { if (auto pointer_type = std::get_if(dst_type.get())) { if (types::equal(*array_type->element_type, *pointer_type->referenced_type)) { auto const arg1_address = node_address(node.arg1); if (arg1_address.base) { builder.lea(*arg1_address.base, arg1_address.offset, reg::rax); } else { push_resolve_global(arg1_address.global); builder.lea_rip_prev(0, 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(builder.code.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) { module.external_function_resolve.push_back({.name = node.name, .instruction_offset = (std::int32_t)builder.code.size()}); builder.mov_read_rip_prev(0, 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(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(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)builder.code.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)builder.code.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)builder.code.size(), node.target}); builder.jump_if_nonzero(0); } template 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(*type) || std::holds_alternative(*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(argument->inferred_type.get()); auto array_type = std::get_if(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(argument->inferred_type.get()); auto array_type = std::get_if(argument->inferred_type.get()); if (struct_type || array_type) { if (auto small_struct = classify_small_struct(lcontext, argument->inferred_type)) { 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) { // Function call node cannot have RIP-relative address auto const address = node_address(it); 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(type.get()); auto array_type = std::get_if(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) { // TODO: call function from a different module? apply_call(it, node, type, [&]{ lcontext.call_resolve.emplace_back(builder.code.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(type.get()); auto array_type = std::get_if(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) 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(*result_type) || std::holds_alternative(*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(&it->instruction)) { stack_position[it] = argument_position[argument->index]; } else if (is_global(it)) { // 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(begin->instruction)) throw std::runtime_error("First IR node of a function must be a label"); auto it = begin; lcontext.nodes[it] = builder.code.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(argument.inferred_type.get()); auto array_type = std::get_if(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] = builder.code.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); } } // Must not be called for global nodes - they use RIP-based addressing with future relocation value_address node_address(ir::node_ref it) { if (is_global(it)) return {.global = it, .offset = 0}; 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 { push_resolve_global(address.global); builder.mov_read_rip_prev(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) { auto const address = node_address(it); builder.mov_write(reg_src, *address.base, address.offset + offset); } else { push_resolve_global(address.global); builder.mov_write_rip_prev(reg_src, 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) { auto const address = node_address(it); 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 { // TODO 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 { // TODO 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)builder.code.size(); reg reg_src; if (src.base) reg_src = *src.base; else { reg_src = find_free_reg(reg::r10, dst.base); push_resolve_global(src.global); 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); push_resolve_global(dst.global); builder.lea_rip_prev(0, 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 push_resolve_global(ir::node_ref global) { module.internal_global_resolve.push_back({.global = global, .instruction_offset = (std::int32_t)builder.code.size()}); } }; } compiled_module compile(ir::compiled_module const & module_in) { compiled_module result; local_context lcontext; { populate_globals_visitor visitor{result, lcontext}; for (auto it = module_in.nodes->begin(); it != module_in.nodes->end(); ++it) std::visit([&](auto const & instruction){ visitor.apply(it, instruction, it->inferred_type); }, it->instruction); } instruction_builder builder{result.code}; for (auto const & function : module_in.functions) { result.functions[function.second.begin] = result.code.size(); compile_visitor visitor{result, lcontext, builder}; visitor.compile(function.first, function.second.begin, function.second.end); } if (module_in.entry_point) result.entry_point = module_in.functions.at(module_in.entry_point).begin; for (auto const & resolve : lcontext.jump_resolve) builder.jump_inject_prev(result.code.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset); for (auto const & resolve : lcontext.cjump_resolve) builder.cjump_inject_prev(result.code.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset); for (auto const & resolve : lcontext.node_resolve) builder.lea_rip_inject_prev(result.code.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset); for (auto const & resolve : lcontext.call_resolve) builder.call_imm_inject_prev(result.code.data() + resolve.offset, lcontext.nodes.at(resolve.target) - resolve.offset); return result; } }