235 lines
7.6 KiB
Go
235 lines
7.6 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: BSD-3-Clause
|
|
|
|
package asm
|
|
|
|
import "fmt"
|
|
|
|
// This file implements VEX (AVX/AVX2) instruction encoding. EVEX (AVX-512)
|
|
// support is a later increment.
|
|
|
|
// vexForm selects how an instruction's operands map onto the VEX.vvvv,
|
|
// ModRM.reg and ModRM.rm fields.
|
|
type vexForm int
|
|
|
|
const (
|
|
// vexNDS3 is the three-operand form `OP src2, src1, dst` (Plan 9 order):
|
|
// ModRM.reg = dst (op2), VEX.vvvv = src1 (op1), ModRM.rm = src2 (op0).
|
|
vexNDS3 vexForm = iota
|
|
// vexRM is the two-operand form `OP src, dst` with no vvvv source:
|
|
// ModRM.reg = dst (op1), ModRM.rm = src (op0), VEX.vvvv = 1111 (unused).
|
|
vexRM
|
|
// vexShiftImm is the immediate-shift form `OP $imm, src, dst`: ModRM.reg =
|
|
// /digit, ModRM.rm = src (op1), VEX.vvvv = dst (op2), imm8 = op0.
|
|
vexShiftImm
|
|
)
|
|
|
|
// vexSpec describes one VEX instruction's encoding parameters.
|
|
type vexSpec struct {
|
|
mapSel int // 1 = 0F, 2 = 0F38, 3 = 0F3A
|
|
opcode byte
|
|
w int // VEX.W (0 for WIG)
|
|
pp int // 0 = none, 1 = 66, 2 = F3, 3 = F2
|
|
opdigit int // ModRM.reg /digit, or -1 when reg is a register
|
|
form vexForm
|
|
}
|
|
|
|
// vexTable maps an upper-case mnemonic to its VEX encoding. It covers the
|
|
// AVX2 instructions used by the go-flac kernels in the three-operand NDS form;
|
|
// it is extended incrementally.
|
|
var vexTable = map[string]vexSpec{
|
|
// VEX.128/256.66.0F.WIG — integer arithmetic / logic / compare.
|
|
"VPADDD": {1, 0xFE, 0, 1, -1, vexNDS3},
|
|
"VPADDQ": {1, 0xD4, 0, 1, -1, vexNDS3},
|
|
"VPSUBD": {1, 0xFA, 0, 1, -1, vexNDS3},
|
|
"VPSUBQ": {1, 0xFB, 0, 1, -1, vexNDS3},
|
|
"VPXOR": {1, 0xEF, 0, 1, -1, vexNDS3},
|
|
"VPOR": {1, 0xEB, 0, 1, -1, vexNDS3},
|
|
"VPAND": {1, 0xDB, 0, 1, -1, vexNDS3},
|
|
"VPANDN": {1, 0xDF, 0, 1, -1, vexNDS3},
|
|
"VPCMPEQD": {1, 0x76, 0, 1, -1, vexNDS3},
|
|
"VPUNPCKLDQ": {1, 0x62, 0, 1, -1, vexNDS3},
|
|
"VPUNPCKHDQ": {1, 0x6A, 0, 1, -1, vexNDS3},
|
|
"VPUNPCKLQDQ": {1, 0x6C, 0, 1, -1, vexNDS3},
|
|
"VPACKSSDW": {1, 0x6B, 0, 1, -1, vexNDS3},
|
|
// VEX.128/256.66.0F38.WIG.
|
|
"VPMULLD": {2, 0x40, 0, 1, -1, vexNDS3},
|
|
"VPMULDQ": {2, 0x28, 0, 1, -1, vexNDS3},
|
|
"VPSHUFB": {2, 0x00, 0, 1, -1, vexNDS3},
|
|
"VPCMPGTQ": {2, 0x37, 0, 1, -1, vexNDS3},
|
|
|
|
// VEX.128/256.66.0F38.WIG — sign/zero extend and broadcast (reg=dst, rm=src,
|
|
// no vvvv).
|
|
"VPMOVSXWD": {2, 0x23, 0, 1, -1, vexRM},
|
|
"VPMOVSXDQ": {2, 0x25, 0, 1, -1, vexRM},
|
|
"VPMOVZXDQ": {2, 0x35, 0, 1, -1, vexRM},
|
|
"VPBROADCASTD": {2, 0x58, 0, 1, -1, vexRM},
|
|
"VPBROADCASTQ": {2, 0x59, 0, 1, -1, vexRM},
|
|
// VEX.128/256.66.0F.WIG — move mask to a GPR (reg=gpr dst, rm=vec src).
|
|
"VPMOVMSKB": {1, 0xD7, 0, 1, -1, vexRM},
|
|
"VMOVMSKPS": {1, 0x50, 0, 0, -1, vexRM}, // no 66 prefix (that would be VMOVMSKPD)
|
|
|
|
// VEX.128/256.66.0F.WIG — immediate shifts (opdigit selects the shift).
|
|
"VPSLLD": {1, 0x72, 0, 1, 6, vexShiftImm},
|
|
"VPSRAD": {1, 0x72, 0, 1, 4, vexShiftImm},
|
|
"VPSRLD": {1, 0x72, 0, 1, 2, vexShiftImm},
|
|
"VPSRLQ": {1, 0x73, 0, 1, 2, vexShiftImm},
|
|
"VPSLLQ": {1, 0x73, 0, 1, 6, vexShiftImm},
|
|
}
|
|
|
|
// isVex reports whether the mnemonic is a VEX-encoded instruction we handle.
|
|
func isVex(mnemUpper string) bool {
|
|
_, ok := vexTable[mnemUpper]
|
|
return ok
|
|
}
|
|
|
|
// encodeVex encodes a VEX instruction with operands in Plan 9 order.
|
|
func (e *enc) encodeVex(mnemUpper string, ops []Operand) error {
|
|
spec := vexTable[mnemUpper]
|
|
switch spec.form {
|
|
case vexNDS3:
|
|
return e.encodeVexNDS3(spec, ops)
|
|
case vexRM:
|
|
return e.encodeVexRM(spec, ops)
|
|
case vexShiftImm:
|
|
return e.encodeVexShiftImm(spec, ops)
|
|
}
|
|
return fmt.Errorf("unhandled VEX form for %s", mnemUpper)
|
|
}
|
|
|
|
// encodeVexNDS3 encodes the three-operand NDS form: OP src2, src1, dst.
|
|
func (e *enc) encodeVexNDS3(spec vexSpec, ops []Operand) error {
|
|
if len(ops) != 3 {
|
|
return fmt.Errorf("VEX NDS instruction expects 3 operands, got %d", len(ops))
|
|
}
|
|
src2, src1, dst := ops[0], ops[1], ops[2]
|
|
|
|
dstReg, ok := dst.(Reg)
|
|
if !ok || !dstReg.isVec() {
|
|
return fmt.Errorf("VEX destination must be a vector register")
|
|
}
|
|
vvvvReg, ok := src1.(Reg)
|
|
if !ok || !vvvvReg.isVec() {
|
|
return fmt.Errorf("VEX vvvv operand must be a vector register")
|
|
}
|
|
|
|
regField := dstReg.idx & 7
|
|
rBit := 0
|
|
if dstReg.idx >= 8 {
|
|
rBit = 1
|
|
}
|
|
vvvvBar := 15 - (vvvvReg.idx & 15)
|
|
return e.emitVexFields(spec, dstReg.vecLenBit(), regField, rBit, vvvvBar, src2)
|
|
}
|
|
|
|
// encodeVexRM encodes the two-operand form: OP src, dst (no vvvv source).
|
|
// ModRM.reg = dst, ModRM.rm = src; the vector length comes from whichever
|
|
// operand is a vector register (the destination for extends/broadcasts, the
|
|
// source for the move-mask instructions whose destination is a GPR).
|
|
func (e *enc) encodeVexRM(spec vexSpec, ops []Operand) error {
|
|
if len(ops) != 2 {
|
|
return fmt.Errorf("VEX two-operand instruction expects 2 operands, got %d", len(ops))
|
|
}
|
|
src, dst := ops[0], ops[1]
|
|
|
|
dstReg, ok := dst.(Reg)
|
|
if !ok {
|
|
return fmt.Errorf("VEX destination must be a register")
|
|
}
|
|
regField := dstReg.idx & 7
|
|
rBit := 0
|
|
if dstReg.idx >= 8 {
|
|
rBit = 1
|
|
}
|
|
|
|
// Vector length: from the destination if it is a vector, otherwise from the
|
|
// source (move-mask instructions have a GPR destination and a vector source).
|
|
l := 0
|
|
if dstReg.isVec() {
|
|
l = dstReg.vecLenBit()
|
|
} else if srcReg, ok := src.(Reg); ok && srcReg.isVec() {
|
|
l = srcReg.vecLenBit()
|
|
}
|
|
|
|
return e.emitVexFields(spec, l, regField, rBit, 0, src) // vvvv unused → vvvvBar=0
|
|
}
|
|
|
|
// encodeVexShiftImm encodes an immediate-shift instruction: OP $imm, src, dst.
|
|
// The destination is carried in VEX.vvvv, the source in ModRM.rm, and the
|
|
// shift kind in the ModRM.reg /digit.
|
|
func (e *enc) encodeVexShiftImm(spec vexSpec, ops []Operand) error {
|
|
if len(ops) != 3 {
|
|
return fmt.Errorf("VEX shift expects 3 operands ($imm, src, dst), got %d", len(ops))
|
|
}
|
|
imm, src, dst := ops[0], ops[1], ops[2]
|
|
immVal, ok := imm.(Imm)
|
|
if !ok {
|
|
return fmt.Errorf("shift count must be an immediate")
|
|
}
|
|
srcReg, ok := src.(Reg)
|
|
if !ok || !srcReg.isVec() {
|
|
return fmt.Errorf("shift source must be a vector register")
|
|
}
|
|
dstReg, ok := dst.(Reg)
|
|
if !ok || !dstReg.isVec() {
|
|
return fmt.Errorf("shift destination must be a vector register")
|
|
}
|
|
|
|
vvvvBar := 15 - (dstReg.idx & 15)
|
|
l := dstReg.vecLenBit()
|
|
rmField := srcReg.idx & 7
|
|
bBit := 0
|
|
if srcReg.idx >= 8 {
|
|
bBit = 1
|
|
}
|
|
modrm := 0xC0 | spec.opdigit<<3 | rmField
|
|
|
|
if spec.mapSel == 1 && bBit == 0 && spec.w == 0 {
|
|
e.out = append(e.out, 0xC5, byte(1<<7|vvvvBar<<3|l<<2|spec.pp))
|
|
} else {
|
|
e.out = append(e.out, 0xC4,
|
|
byte(1<<7|1<<6|(1-bBit)<<5|spec.mapSel),
|
|
byte(spec.w<<7|vvvvBar<<3|l<<2|spec.pp))
|
|
}
|
|
e.out = append(e.out, spec.opcode, byte(modrm), byte(int8(immVal)))
|
|
return nil
|
|
}
|
|
|
|
// emitVexFields emits the VEX prefix, opcode, ModR/M, SIB and displacement for
|
|
// the given precomputed fields. It is shared by the NDS and RM forms.
|
|
func (e *enc) emitVexFields(spec vexSpec, l, regField, rBit, vvvvBar int, rm Operand) error {
|
|
var modrm, sib int
|
|
var disp []byte
|
|
var xBit, bBit int
|
|
switch r := rm.(type) {
|
|
case Reg:
|
|
modrm = 0xC0 | regField<<3 | (r.idx & 7)
|
|
sib = -1
|
|
if r.idx >= 8 {
|
|
bBit = 1
|
|
}
|
|
case Mem:
|
|
var err error
|
|
modrm, sib, disp, xBit, bBit, err = memComponents(regField, r)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
default:
|
|
return fmt.Errorf("invalid VEX r/m operand")
|
|
}
|
|
|
|
if spec.mapSel == 1 && xBit == 0 && bBit == 0 && spec.w == 0 {
|
|
e.out = append(e.out, 0xC5, byte((1-rBit)<<7|vvvvBar<<3|l<<2|spec.pp))
|
|
} else {
|
|
e.out = append(e.out, 0xC4,
|
|
byte((1-rBit)<<7|(1-xBit)<<6|(1-bBit)<<5|spec.mapSel),
|
|
byte(spec.w<<7|vvvvBar<<3|l<<2|spec.pp))
|
|
}
|
|
e.out = append(e.out, spec.opcode, byte(modrm))
|
|
if sib >= 0 {
|
|
e.out = append(e.out, byte(sib))
|
|
}
|
|
e.out = append(e.out, disp...)
|
|
return nil
|
|
}
|