Files
jak-project/goalc/emitter/Instruction.h
T
Tyler Wilding 76941379e9 goalc: Finish converting all load/store instructions to ARM64 (#4334)
All that remains (8 instructions) are division, and NEON instructions
that require me to convert the x86 control byte to NEON `TBL` values.
Those handful of instructions can be done later while doing the next
steps (finally something more interesting than just encoding
instructions).
2026-07-03 19:40:29 -04:00

1282 lines
28 KiB
C++

#pragma once
#include <cstring>
#include <span>
#include <variant>
#include "common/common_types.h"
#include "common/util/Assert.h"
namespace emitter {
/*!
* A high-level description of a opcode. It can emit itself.
*/
template <typename InstructionType>
struct InstructionImpl {
/*!
* Emit into a buffer and return how many bytes written (can be zero)
*/
u8 emit(u8* buffer) const { return static_cast<const InstructionType*>(this)->emit(buffer); }
// TODO - the below might only be relevant for X86, in which case
// they can eventually leave this parent type
// and at that point, things can likely be simplified
//
// For now, just trying to make things compile / work
u8 length() const { return static_cast<const InstructionType*>(this)->length(); }
int get_imm_size() const { return static_cast<const InstructionType*>(this)->get_imm_size(); }
int get_disp_size() const { return static_cast<const InstructionType*>(this)->get_disp_size(); }
int offset_of_imm() const { return static_cast<const InstructionType*>(this)->offset_of_imm(); }
int offset_of_disp() const { return static_cast<const InstructionType*>(this)->offset_of_disp(); }
};
namespace ARM64 {
struct Field {
u32 bits;
constexpr explicit Field(u32 v) : bits(v) {}
};
constexpr u32 Base(u32 value, u32 width) {
return value << (32 - width);
}
// TODO - consider passing in the instruction name to make debugging easier when an assertion is
// hit
// TODO NOW - fix below
constexpr u64 pow2(u64 n) {
return 1ull << n;
}
constexpr s64 pow2s(u64 n) {
return 1ull << n;
}
constexpr Field Hw(u32 x) {
ASSERT(x >= 0 && x <= (4 - 1));
return Field{(x & 4) << 21};
}
constexpr Field Sh(u32 x) {
ASSERT(x >= 0 && x <= (2 - 1));
return Field{(x & 1) << 22};
}
constexpr Field Shift(u32 x) {
ASSERT(x >= 0 && x <= (4 - 1));
return Field{(x & 2) << 22};
}
constexpr Field Rd(u32 x) {
ASSERT(x >= 0 && x <= (32 - 1));
return Field{(x & 31) << 0};
}
constexpr Field Rt(u32 x) {
ASSERT(x >= 0 && x <= (32 - 1));
return Field{(x & 31) << 0};
}
constexpr Field Rn(u32 x) {
ASSERT(x >= 0 && x <= (32 - 1));
return Field{(x & 31) << 5};
}
constexpr Field Rm(u32 x) {
ASSERT(x >= 0 && x <= (32 - 1));
return Field{(x & 31) << 16};
}
constexpr Field Imm4(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 4) - 1));
return Field{(x & 0b111111) << 11};
}
constexpr Field Imm6(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 6)));
return Field{(x & 0b111111) << 10};
}
constexpr Field Imm9s(s32 x) {
ASSERT(x >= (pow2s(9 - 1) * -1) && x <= (pow2s(9 - 1) - 1));
return Field{(static_cast<u32>(x) & 0b111111111) << 12};
}
constexpr Field Imm12(u32 x) {
ASSERT(x >= 0 && x <= (pow2(12) - 1));
return Field{(static_cast<u32>(x) & 0b111111111111) << 10};
}
constexpr Field Imm16(u32 x) {
ASSERT(x >= 0 && x <= (pow2(16) - 1));
return Field{static_cast<u32>((x & (pow2(16) - 1)) << 16)};
}
constexpr Field Imm26(u32 x) {
ASSERT(x >= 0 && x <= (67108864 - 1));
return Field{(static_cast<uint32_t>(x) & 0b11111111111111111111111111) << 0};
}
constexpr Field Imm19(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 19) - 1));
return Field{(static_cast<uint32_t>(x) & 0b1111111111111111111) << 5};
}
constexpr Field Immlo(u32 x) {
ASSERT(x >= 0 && x <= (pow2(2) - 1));
return Field{(static_cast<u32>(x) & 0b11) << 29};
}
constexpr Field Immhi(u32 x) {
ASSERT(x >= 0 && x <= (pow2(19) - 1));
return Field{(static_cast<u32>(x) & 0b1111111111111111111) << 5};
}
constexpr Field Imms(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 6) - 1));
return Field{(static_cast<uint32_t>(x) & 0b111111) << 10};
}
constexpr Field Immr(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 6) - 1));
return Field{(static_cast<uint32_t>(x) & 0b111111) << 16};
}
constexpr Field Immh(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 4) - 1));
return Field{(static_cast<uint32_t>(x) & 0b111111) << 19};
}
constexpr Field Immb(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 3) - 1));
return Field{(static_cast<uint32_t>(x) & 0b111111) << 16};
}
constexpr Field Cond(u32 x) {
ASSERT(x >= 0 && x <= ((2 ^ 4) - 1));
return Field{(static_cast<uint32_t>(x) & 0b1111) << 0};
}
} // namespace ARM64
struct InstructionARM64 : InstructionImpl<InstructionARM64> {
// The ARM instruction stream is a sequence of word-aligned words.
// Each ARM instruction is a single 32-bit word in that stream.
//
// Some x86 instructions are not possible to represent in ARM in a single instruction
// however, in order to not have to overhaul things at the IR level,
// it feels preferably to instead allow an instruction to emit multiple instructions if needed
//
// To do so, the instruction can optionally include multiple encodings
// all of which are emitted at once.
static constexpr int kMaxInstrs = 64;
u32 encodings[kMaxInstrs]{};
u8 count = 0;
InstructionARM64() = delete;
// --- single instruction ---
template <typename... Fs>
constexpr InstructionARM64(uint32_t base, Fs... fields) {
static_assert((std::is_same_v<Fs, emitter::ARM64::Field> && ...));
encodings[0] = (base | ... | fields.bits);
count = 1;
}
// --- multi instruction (variadic) ---
template <typename... Instrs>
constexpr InstructionARM64(const Instrs&... instrs)
requires(std::is_same_v<Instrs, InstructionARM64> && ...)
{
u8 idx = 0;
auto append = [&](const InstructionARM64& i) {
for (uint8_t j = 0; j < i.count; ++j) {
encodings[idx++] = i.encodings[j];
}
};
(append(instrs), ...);
count = idx;
}
InstructionARM64(std::span<const InstructionARM64> instrs) {
u8 idx = 0;
for (const auto& i : instrs) {
for (uint8_t j = 0; j < i.count; ++j) {
encodings[idx++] = i.encodings[j];
}
}
count = idx;
}
uint8_t emit(uint8_t* buffer) const {
if (count == 1 && encodings[0] == 0) {
return 0;
}
memcpy(buffer, encodings, count * 4);
return count * 4;
}
uint8_t length() const {
if (count == 1 && encodings[0] == 0) {
return 0;
}
return count * 4;
}
// TODO ARM - all placeholders, no idea if this is even relevant, if not, get rid of it all
int get_imm_size() const { return 0; }
int offset_of_imm() const { return 0; }
int offset_of_disp() const { return 0; }
int get_disp_size() const { return 0; }
};
/*!
* The ModRM byte
*/
struct ModRM {
uint8_t mod;
uint8_t reg_op;
uint8_t rm;
uint8_t operator()() const { return (mod << 6) | (reg_op << 3) | (rm << 0); }
};
/*!
* The SIB Byte
*/
struct SIB {
uint8_t scale, index, base;
uint8_t operator()() const { return (scale << 6) | (index << 3) | (base << 0); }
};
/*!
* An Immediate (either imm or disp)
*/
struct Imm {
Imm() = default;
Imm(uint8_t sz, uint64_t v) : size(sz), value(v) {}
uint8_t size;
union {
uint64_t value;
uint8_t v_arr[8];
};
};
/*!
* The REX prefix byte
*/
struct REX {
explicit REX(bool w = false, bool r = false, bool x = false, bool b = false)
: W(w), R(r), X(x), B(b) {}
// W - 64-bit operands
// R - reg extension
// X - SIB i extnsion
// B - other extension
bool W, R, X, B;
uint8_t operator()() const { return (1 << 6) | (W << 3) | (R << 2) | (X << 1) | (B << 0); }
};
enum class VexPrefix : u8 { P_NONE = 0, P_66 = 1, P_F3 = 2, P_F2 = 3 };
/*!
* The "VEX" 3-byte format for AVX instructions
*/
struct VEX3 {
bool W, R, X, B;
enum class LeadingBytes : u8 { P_INVALID = 0, P_0F = 1, P_0F_38 = 2, P_0F_3A = 3 } leading_bytes;
u8 reg_id;
VexPrefix prefix;
bool L;
u8 emit(u8 byte) const {
if (byte == 0) {
return 0b11000100;
} else if (byte == 1) {
u8 result = 0;
result |= ((!R) << 7);
result |= ((!X) << 6);
result |= ((!B) << 5);
result |= (0b11111 & u8(leading_bytes));
return result;
} else if (byte == 2) {
u8 result = 0;
result |= (W << 7); // this may be inverted?
result |= ((~reg_id) & 0b1111) << 3;
result |= (L << 2);
result |= (u8(prefix) & 0b11);
return result;
} else {
ASSERT(false);
return -1;
}
}
VEX3(bool w,
bool r,
bool x,
bool b,
LeadingBytes _leading_bytes,
u8 _reg_id = 0,
VexPrefix _prefix = VexPrefix::P_NONE,
bool l = false)
: W(w),
R(r),
X(x),
B(b),
leading_bytes(_leading_bytes),
reg_id(_reg_id),
prefix(_prefix),
L(l) {}
};
struct VEX2 {
bool R;
u8 reg_id;
VexPrefix prefix;
bool L;
u8 emit(u8 byte) const {
if (byte == 0) {
return 0b11000101;
} else if (byte == 1) {
u8 result = 0;
result |= ((!R) << 7);
result |= ((~reg_id) & 0b1111) << 3;
result |= (L << 2);
result |= (u8(prefix) & 0b11);
return result;
} else {
ASSERT(false);
return -1;
}
}
VEX2(bool r, u8 _reg_id = 0, VexPrefix _prefix = VexPrefix::P_NONE, bool l = false)
: R(r), reg_id(_reg_id), prefix(_prefix), L(l) {}
};
struct InstructionX86 : InstructionImpl<InstructionX86> {
enum Flags {
kOp2Set = (1 << 0),
kOp3Set = (1 << 1),
kIsNull = (1 << 2),
kSetRex = (1 << 3),
kSetModrm = (1 << 4),
kSetSib = (1 << 5),
kSetDispImm = (1 << 6),
kSetImm = (1 << 7),
};
InstructionX86(u8 opcode) : op(opcode) {}
u8 op;
u8 m_flags = 0;
u8 op2;
u8 op3;
u8 n_vex = 0;
u8 vex[3] = {0, 0, 0};
// the rex byte
u8 m_rex = 0;
// the modrm byte
u8 m_modrm = 0;
// the sib byte
u8 m_sib = 0;
// the displacement
Imm disp;
// the immediate
Imm imm;
/*!
* Move opcode byte 0 to before the rex prefix.
*/
void swap_op0_rex() {
if (!(m_flags & kSetRex))
return;
auto temp = op;
op = m_rex;
m_rex = temp;
}
void set(REX r) {
m_rex = r();
m_flags |= kSetRex;
}
void set(ModRM modrm) {
m_modrm = modrm();
m_flags |= kSetModrm;
}
void set(SIB sib) {
m_sib = sib();
m_flags |= kSetSib;
}
void set(VEX3 vex3) {
n_vex = 3;
for (int i = 0; i < n_vex; i++) {
vex[i] = vex3.emit(i);
}
}
void set(VEX2 vex2) {
n_vex = 2;
for (int i = 0; i < n_vex; i++) {
vex[i] = vex2.emit(i);
}
}
void set_disp(Imm i) {
disp = i;
m_flags |= kSetDispImm;
}
void set(Imm i) {
imm = i;
m_flags |= kSetImm;
}
void set_op2(uint8_t b) {
m_flags |= kOp2Set;
op2 = b;
}
void set_op3(uint8_t b) {
m_flags |= kOp3Set;
op3 = b;
}
int get_imm_size() const {
if (m_flags & kSetImm) {
return imm.size;
} else {
return 0;
}
}
int get_disp_size() const {
if (m_flags & kSetDispImm) {
return disp.size;
} else {
return 0;
}
}
/*!
* Set modrm and rex as needed for two regs.
*/
void set_modrm_and_rex(uint8_t reg, uint8_t rm, uint8_t mod, bool rex_w = false) {
bool rex_b = false, rex_r = false;
if (rm >= 8) {
rm -= 8;
rex_b = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = mod;
modrm.reg_op = reg;
modrm.rm = rm;
set(modrm);
if (rex_b || rex_w || rex_r) {
set(REX(rex_w, rex_r, false, rex_b));
}
}
void set_vex_modrm_and_rex(uint8_t reg,
uint8_t rm,
VEX3::LeadingBytes lb,
uint8_t vex_reg = 0,
bool rex_w = false,
VexPrefix prefix = VexPrefix::P_NONE) {
bool rex_b = false, rex_r = false;
if (rm >= 8) {
rm -= 8;
rex_b = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 3;
modrm.reg_op = reg;
modrm.rm = rm;
set(modrm);
if (rex_b || rex_w || lb != VEX3::LeadingBytes::P_0F) {
// need three byte version
set(VEX3(rex_w, rex_r, false, rex_b, lb, vex_reg, prefix));
} else {
ASSERT(lb == VEX3::LeadingBytes::P_0F); // vex2 implies 0x0f
ASSERT(!rex_b);
ASSERT(!rex_w);
set(VEX2(rex_r, vex_reg, prefix));
}
}
/*!
* Set VEX prefix for REX as needed for two registers.
*/
void set_vex_modrm_and_rex(uint8_t reg,
uint8_t rm,
uint8_t mod,
VEX3::LeadingBytes lb,
bool rex_w = false) {
bool rex_b = false;
bool rex_r = false;
if (rm >= 8) {
rm -= 8;
rex_b = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = mod;
modrm.reg_op = reg;
modrm.rm = rm;
set(modrm);
if (rex_b || rex_w || lb != VEX3::LeadingBytes::P_0F) {
// need three byte version
set(VEX3(rex_w, rex_r, false, rex_b, lb));
} else {
// can get away with two byte version
ASSERT(lb == VEX3::LeadingBytes::P_0F); // vex2 implies 0x0f
ASSERT(!rex_b);
ASSERT(!rex_w);
set(VEX2(rex_r));
}
}
void set_modrm_and_rex_for_reg_plus_reg_plus_s8(uint8_t reg,
uint8_t addr1,
uint8_t addr2,
s8 offset,
bool rex_w) {
bool rex_b = false, rex_r = false, rex_x = false;
bool addr1_ext = false;
bool addr2_ext = false;
if (addr1 >= 8) {
addr1 -= 8;
addr1_ext = true;
}
if (addr2 >= 8) {
addr2 -= 8;
addr2_ext = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 1; // no disp
modrm.rm = 4; // sib!
modrm.reg_op = reg;
SIB sib;
sib.scale = 0;
Imm imm2(1, offset);
// default addr1 in index
if (addr1 == 4) {
sib.index = addr2;
sib.base = addr1;
rex_x = addr2_ext;
rex_b = addr1_ext;
} else {
// addr1 in index
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
}
ASSERT(sib.index != 4);
if (rex_b || rex_w || rex_r || rex_x) {
set(REX(rex_w, rex_r, rex_x, rex_b));
}
set(modrm);
set(sib);
set_disp(imm2);
}
void set_vex_modrm_and_rex_for_reg_plus_reg_plus_s8(uint8_t reg,
uint8_t addr1,
uint8_t addr2,
s8 offset,
VEX3::LeadingBytes lb,
bool rex_w) {
bool rex_b = false, rex_r = false, rex_x = false;
bool addr1_ext = false;
bool addr2_ext = false;
if (addr1 >= 8) {
addr1 -= 8;
addr1_ext = true;
}
if (addr2 >= 8) {
addr2 -= 8;
addr2_ext = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 1; // no disp
modrm.rm = 4; // sib!
modrm.reg_op = reg;
SIB sib;
sib.scale = 0;
Imm imm2(1, offset);
// default addr1 in index
if (addr1 == 4) {
sib.index = addr2;
sib.base = addr1;
rex_x = addr2_ext;
rex_b = addr1_ext;
} else {
// addr1 in index
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
}
ASSERT(sib.index != 4);
if (rex_b || rex_w || rex_x || lb != VEX3::LeadingBytes::P_0F) {
// need three byte version
set(VEX3(rex_w, rex_r, rex_x, rex_b, lb));
} else {
ASSERT(lb == VEX3::LeadingBytes::P_0F); // vex2 implies 0x0f
ASSERT(!rex_b);
ASSERT(!rex_w);
ASSERT(!rex_x);
set(VEX2(rex_r));
}
set(modrm);
set(sib);
set_disp(imm2);
}
void set_modrm_and_rex_for_reg_plus_reg_plus_s32(uint8_t reg,
uint8_t addr1,
uint8_t addr2,
s32 offset,
bool rex_w) {
bool rex_b = false, rex_r = false, rex_x = false;
bool addr1_ext = false;
bool addr2_ext = false;
if (addr1 >= 8) {
addr1 -= 8;
addr1_ext = true;
}
if (addr2 >= 8) {
addr2 -= 8;
addr2_ext = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 2; // no disp
modrm.rm = 4; // sib!
modrm.reg_op = reg;
SIB sib;
sib.scale = 0;
Imm imm2(4, offset);
// default addr1 in index
if (addr1 == 4) {
sib.index = addr2;
sib.base = addr1;
rex_x = addr2_ext;
rex_b = addr1_ext;
} else {
// addr1 in index
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
}
ASSERT(sib.index != 4);
if (rex_b || rex_w || rex_r || rex_x) {
set(REX(rex_w, rex_r, rex_x, rex_b));
}
set(modrm);
set(sib);
set_disp(imm2);
}
void set_vex_modrm_and_rex_for_reg_plus_reg_plus_s32(uint8_t reg,
uint8_t addr1,
uint8_t addr2,
s32 offset,
VEX3::LeadingBytes lb,
bool rex_w) {
bool rex_b = false, rex_r = false, rex_x = false;
bool addr1_ext = false;
bool addr2_ext = false;
if (addr1 >= 8) {
addr1 -= 8;
addr1_ext = true;
}
if (addr2 >= 8) {
addr2 -= 8;
addr2_ext = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 2; // no disp
modrm.rm = 4; // sib!
modrm.reg_op = reg;
SIB sib;
sib.scale = 0;
Imm imm2(4, offset);
// default addr1 in index
if (addr1 == 4) {
sib.index = addr2;
sib.base = addr1;
rex_x = addr2_ext;
rex_b = addr1_ext;
} else {
// addr1 in index
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
}
ASSERT(sib.index != 4);
if (rex_b || rex_w || rex_x || lb != VEX3::LeadingBytes::P_0F) {
// need three byte version
set(VEX3(rex_w, rex_r, rex_x, rex_b, lb));
} else {
ASSERT(lb == VEX3::LeadingBytes::P_0F); // vex2 implies 0x0f
ASSERT(!rex_b);
ASSERT(!rex_w);
ASSERT(!rex_x);
set(VEX2(rex_r));
}
set(modrm);
set(sib);
set_disp(imm2);
}
void set_modrm_and_rex_for_reg_plus_reg_addr(uint8_t reg,
uint8_t addr1,
uint8_t addr2,
bool rex_w = false,
bool rex_always = false) {
bool rex_b = false, rex_r = false, rex_x = false;
bool addr1_ext = false;
bool addr2_ext = false;
if (addr1 >= 8) {
addr1 -= 8;
addr1_ext = true;
}
if (addr2 >= 8) {
addr2 -= 8;
addr2_ext = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 0; // no disp
modrm.rm = 4; // sib!
modrm.reg_op = reg;
SIB sib;
sib.scale = 0;
if (addr1 == 5 && addr2 == 5) {
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
modrm.mod = 1;
set_disp(Imm(1, 0));
} else {
// default addr1 in index
bool flipped = (addr1 == 4) || (addr2 == 5);
if (flipped) {
sib.index = addr2;
sib.base = addr1;
rex_x = addr2_ext;
rex_b = addr1_ext;
} else {
// addr1 in index
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
}
ASSERT(sib.base != 5);
ASSERT(sib.index != 4);
}
if (rex_b || rex_w || rex_r || rex_x || rex_always) {
set(REX(rex_w, rex_r, rex_x, rex_b));
}
set(modrm);
set(sib);
}
void set_vex_modrm_and_rex_for_reg_plus_reg_addr(uint8_t reg,
uint8_t addr1,
uint8_t addr2,
VEX3::LeadingBytes lb,
bool rex_w = false) {
bool rex_b = false, rex_r = false, rex_x = false;
bool addr1_ext = false;
bool addr2_ext = false;
if (addr1 >= 8) {
addr1 -= 8;
addr1_ext = true;
}
if (addr2 >= 8) {
addr2 -= 8;
addr2_ext = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 0; // no disp
modrm.rm = 4; // sib!
modrm.reg_op = reg;
SIB sib;
sib.scale = 0;
if (addr1 == 5 && addr2 == 5) {
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
modrm.mod = 1;
set_disp(Imm(1, 0));
} else {
// default addr1 in index
bool flipped = (addr1 == 4) || (addr2 == 5);
if (flipped) {
sib.index = addr2;
sib.base = addr1;
rex_x = addr2_ext;
rex_b = addr1_ext;
} else {
// addr1 in index
sib.index = addr1;
sib.base = addr2;
rex_x = addr1_ext;
rex_b = addr2_ext;
}
ASSERT(sib.base != 5);
ASSERT(sib.index != 4);
}
if (rex_b || rex_w || rex_x || lb != VEX3::LeadingBytes::P_0F) {
// need three byte version
set(VEX3(rex_w, rex_r, rex_x, rex_b, lb));
} else {
ASSERT(lb == VEX3::LeadingBytes::P_0F); // vex2 implies 0x0f
ASSERT(!rex_b);
ASSERT(!rex_w);
ASSERT(!rex_x);
set(VEX2(rex_r));
}
set(modrm);
set(sib);
}
/*!
* Set modrm and rex as needed for two regs for an addressing mode.
* Will set SIB if R12 or RSP indexing is used.
*/
void set_modrm_and_rex_for_reg_addr(uint8_t reg, uint8_t rm, bool rex_w = false) {
bool rex_b = false, rex_r = false;
if (rm >= 8) {
rm -= 8;
rex_b = true;
}
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 0;
modrm.reg_op = reg;
modrm.rm = rm;
if (rm == 4) {
SIB sib;
sib.scale = 0;
sib.base = 4;
sib.index = 4;
set(sib);
}
if (rm == 5) {
modrm.mod = 1; // 1 byte imm
set_disp(Imm(1, 0));
}
set(modrm);
if (rex_b || rex_w || rex_r) {
set(REX(rex_w, rex_r, false, rex_b));
}
}
void set_modrm_and_rex_for_rip_plus_s32(uint8_t reg, s32 offset, bool rex_w = false) {
bool rex_r = false;
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 0;
modrm.reg_op = reg;
modrm.rm = 5; // use the RIP addressing mode
set(modrm);
if (rex_r || rex_w) {
set(REX(rex_w, rex_r, false, false));
}
set_disp(Imm(4, offset));
}
void add_rex() {
if (!(m_flags & kSetRex)) {
set(REX());
}
}
void set_vex_modrm_and_rex_for_rip_plus_s32(uint8_t reg,
s32 offset,
VEX3::LeadingBytes lb = VEX3::LeadingBytes::P_0F,
bool rex_w = false) {
bool rex_r = false;
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
ModRM modrm;
modrm.mod = 0;
modrm.reg_op = reg;
modrm.rm = 5; // use the RIP addressing mode
set(modrm);
if (rex_w || lb != VEX3::LeadingBytes::P_0F) {
// need three byte version
set(VEX3(rex_w, rex_r, false, false, lb));
} else {
ASSERT(lb == VEX3::LeadingBytes::P_0F); // vex2 implies 0x0f
ASSERT(!rex_w);
set(VEX2(rex_r));
}
set_disp(Imm(4, offset));
}
/*!
* Set up modrm and rex for the commonly used immediate displacement indexing mode.
*/
void set_modrm_rex_sib_for_reg_reg_disp(uint8_t reg, uint8_t mod, uint8_t rm, bool rex_w) {
ModRM modrm;
bool rex_r = false;
if (reg >= 8) {
reg -= 8;
rex_r = true;
}
modrm.reg_op = reg;
modrm.mod = mod;
modrm.rm = 4; // use sib
SIB sib;
sib.scale = 0;
sib.index = 4;
bool rex_b = false;
if (rm >= 8) {
rex_b = true;
rm -= 8;
}
sib.base = rm;
set(modrm);
set(sib);
if (rex_r || rex_w || rex_b) {
set(REX(rex_w, rex_r, false, rex_b));
}
}
/*!
* Get the position of the disp immediate relative to the start of the instruction
*/
int offset_of_disp() const {
if (m_flags & kIsNull)
return 0;
ASSERT(m_flags & kSetDispImm);
int offset = 0;
offset += n_vex;
if (m_flags & kSetRex)
offset++;
offset++; // opcode
if (m_flags & kOp2Set)
offset++;
if (m_flags & kOp3Set)
offset++;
if (m_flags & kSetModrm)
offset++;
if (m_flags & kSetSib)
offset++;
return offset;
}
/*!
* Get the position of the imm immediate relative to the start of the instruction
*/
int offset_of_imm() const {
if (m_flags & kIsNull)
return 0;
ASSERT(m_flags & kSetImm);
int offset = 0;
offset += n_vex;
if (m_flags & kSetRex)
offset++;
offset++; // opcode
if (m_flags & kOp2Set)
offset++;
if (m_flags & kOp3Set)
offset++;
if (m_flags & kSetModrm)
offset++;
if (m_flags & kSetSib)
offset++;
if (m_flags & kSetDispImm)
offset += disp.size;
return offset;
}
uint8_t emit(uint8_t* buffer) const {
if (m_flags & kIsNull)
return 0;
uint8_t count = 0;
for (int i = 0; i < n_vex; i++) {
buffer[count++] = vex[i];
}
if (m_flags & kSetRex) {
buffer[count++] = m_rex;
}
buffer[count++] = op;
if (m_flags & kOp2Set) {
buffer[count++] = op2;
}
if (m_flags & kOp3Set) {
buffer[count++] = op3;
}
if (m_flags & kSetModrm) {
buffer[count++] = m_modrm;
}
if (m_flags & kSetSib) {
buffer[count++] = m_sib;
}
if (m_flags & kSetDispImm) {
for (int i = 0; i < disp.size; i++) {
buffer[count++] = disp.v_arr[i];
}
}
if (m_flags & kSetImm) {
for (int i = 0; i < imm.size; i++) {
buffer[count++] = imm.v_arr[i];
}
}
return count;
}
uint8_t length() const {
if (m_flags & kIsNull)
return 0;
uint8_t count = 0;
count += n_vex;
if (m_flags & kSetRex) {
count++;
}
count++;
if (m_flags & kOp2Set) {
count++;
}
if (m_flags & kOp3Set) {
count++;
}
if (m_flags & kSetModrm) {
count++;
}
if (m_flags & kSetSib) {
count++;
}
if (m_flags & kSetDispImm) {
for (int i = 0; i < disp.size; i++) {
count++;
}
}
if (m_flags & kSetImm) {
for (int i = 0; i < imm.size; i++) {
count++;
}
}
return count;
}
};
class Instruction {
public:
using Variant = std::variant<InstructionX86, InstructionARM64>;
Variant instr;
Instruction() = delete;
template <typename T>
Instruction(T v) : instr(std::move(v)) {}
u8 emit(u8* buffer) const {
return std::visit([&](auto const& i) { return i.emit(buffer); }, instr);
}
u8 length() const {
return std::visit([](auto const& i) { return i.length(); }, instr);
}
int get_imm_size() const {
return std::visit([](auto const& i) { return i.get_imm_size(); }, instr);
}
int get_disp_size() const {
return std::visit([](auto const& i) { return i.get_disp_size(); }, instr);
}
int offset_of_imm() const {
return std::visit([](auto const& i) { return i.offset_of_imm(); }, instr);
}
int offset_of_disp() const {
return std::visit([](auto const& i) { return i.offset_of_disp(); }, instr);
}
};
} // namespace emitter