feat(asm): byte-identical go-flac AVX2 assembly with scalar families and jump relaxation

Assisted-by: Qwen 3.8 Max Preview
This commit is contained in:
2026-07-08 12:51:35 +02:00
parent 39870f91f6
commit a82f575aee
12 changed files with 662 additions and 73 deletions
+149 -35
View File
@@ -13,9 +13,10 @@ import (
// Assemble encodes the body of a TEXT function into x86-64 machine code,
// resolving local labels to relative jump offsets and translating the FP/SP
// pseudo-registers onto the hardware stack pointer (matching the Go
// assembler's default frame-pointer behaviour). Jumps always use the 32-bit
// relative form so instruction sizes are fixed and offsets resolve in a single
// layout pass.
// assembler's default frame-pointer behaviour). Jumps start in the short
// (rel8) form and expand to rel32 when the settled displacement does not fit;
// sizes only grow, so the layout reaches a fixed point in a few passes. CALL
// has no short form and is always rel32.
//
// Supported operands: registers, memory (real base register), immediates,
// FP/SP frame-relative operands, and local-label jumps. SB (global symbol)
@@ -23,35 +24,74 @@ import (
// integer and shuffle/extract/permute/move set is in.
func Assemble(t *ast.Text) ([]byte, map[string]int, error) {
fi := computeFrame(t)
chain := jumpChain(t)
resolve := func(name string) string {
if r, ok := chain[name]; ok {
return r
}
return name
}
// Pass 1: lay out instructions (including prologue/epilogue) to fix label
// offsets.
offsets := map[string]int{}
// Layout: iterate jump sizes to a fixed point.
long := make([]bool, len(t.Body))
sizes := make([]int, len(t.Body))
pos := len(fi.prologue)
for i, stmt := range t.Body {
switch s := stmt.(type) {
case *ast.Label:
offsets[s.Name.Text] = pos
case *ast.Instr:
sz, err := instrSize(s, fi)
if err != nil {
return nil, nil, fmt.Errorf("%s: %w", s.Mnemonic.Text, err)
offsets := map[string]int{}
pcs := make([]int, len(t.Body))
for {
pos := len(fi.prologue)
for i, stmt := range t.Body {
switch s := stmt.(type) {
case *ast.Label:
offsets[s.Name.Text] = pos
case *ast.Instr:
sz, err := instrSize(s, fi, long[i])
if err != nil {
return nil, nil, fmt.Errorf("%s: %w", s.Mnemonic.Text, err)
}
sizes[i] = sz
pcs[i] = pos
pos += sz
}
sizes[i] = sz
pos += sz
}
// Expand any short jump whose displacement no longer fits rel8.
changed := false
for i, stmt := range t.Body {
s, ok := stmt.(*ast.Instr)
if !ok {
continue
}
mnem := strings.ToUpper(s.Mnemonic.Text)
if !isJumpMnemonic(mnem) || mnem == "CALL" || long[i] {
continue
}
name, ok := labelName(s.Operands[0])
if !ok {
continue // reported during emission
}
target, ok := offsets[resolve(name)]
if !ok {
continue // reported during emission
}
rel := int64(target - (pcs[i] + jumpSize(mnem, false)))
if !fits8(rel) {
long[i] = true
changed = true
}
}
if !changed {
break
}
}
// Pass 2: emit.
out := append([]byte(nil), fi.prologue...)
pos = len(fi.prologue)
pos := len(fi.prologue)
for i, stmt := range t.Body {
s, ok := stmt.(*ast.Instr)
if !ok {
continue
}
code, err := encodeInstr(s, pos, offsets, fi)
code, err := encodeInstr(s, pos, offsets, fi, long[i], resolve)
if err != nil {
return nil, nil, fmt.Errorf("%s: %w", s.Mnemonic.Text, err)
}
@@ -64,6 +104,58 @@ func Assemble(t *ast.Text) ([]byte, map[string]int, error) {
return out, offsets, nil
}
// jumpChain precomputes jump-to-jump folding: a label whose first instruction
// is an unconditional local jump redirects its own jumpers to the ultimate
// target. The Go toolchain chases exactly these chains (the linker's xfol
// pass) before it encodes branches, so matching its bytes requires the same
// redirection.
func jumpChain(t *ast.Text) map[string]string {
// label → the target of its leading unconditional local JMP, if any.
leadsTo := map[string]string{}
for i, stmt := range t.Body {
l, ok := stmt.(*ast.Label)
if !ok {
continue
}
// Stacked labels share an address: skip to the first instruction.
j := i + 1
for j < len(t.Body) {
if _, isLabel := t.Body[j].(*ast.Label); !isLabel {
break
}
j++
}
if j >= len(t.Body) {
continue
}
in, ok := t.Body[j].(*ast.Instr)
if !ok || strings.ToUpper(in.Mnemonic.Text) != "JMP" || len(in.Operands) != 1 {
continue
}
if name, ok := labelName(in.Operands[0]); ok {
leadsTo[l.Name.Text] = name
}
}
// Chase each chain to its end, guarding against cycles.
chain := map[string]string{}
for name := range leadsTo {
visited := map[string]bool{name: true}
cur := name
for {
next, ok := leadsTo[cur]
if !ok || visited[next] {
break
}
visited[next] = true
cur = next
}
if cur != name {
chain[name] = cur
}
}
return chain
}
// frameInfo carries the frame layout derived from the TEXT directive.
type frameInfo struct {
size int // local frame size ($framesize)
@@ -119,15 +211,15 @@ func addSP(size int) []byte { // ADDQ $size, SP
return append([]byte{0x48, 0x81, 0xC4}, le32(int64(size))...)
}
// instrSize returns the encoded length of an instruction (pass 1). encodeInstr
// already includes the epilogue for a RET in a frame-pointer function; jumps use
// a fixed rel32 size (no epilogue).
func instrSize(s *ast.Instr, fi frameInfo) (int, error) {
// instrSize returns the encoded length of an instruction (layout pass).
// encodeInstr already includes the epilogue for a RET in a frame-pointer
// function; jumps use their short or long form (never an epilogue).
func instrSize(s *ast.Instr, fi frameInfo, long bool) (int, error) {
mnem := strings.ToUpper(s.Mnemonic.Text)
if isJumpMnemonic(mnem) {
return jumpSize(mnem), nil
return jumpSize(mnem, long), nil
}
code, err := encodeInstr(s, 0, nil, fi)
code, err := encodeInstr(s, 0, nil, fi, false, nil)
if err != nil {
return 0, err
}
@@ -142,18 +234,27 @@ func isJumpMnemonic(mnem string) bool {
return ok
}
// jumpSize returns the fixed length of a rel32 jump instruction.
func jumpSize(mnem string) int {
if mnem == "JMP" || mnem == "CALL" {
// jumpSize returns the length of a jump instruction in the requested form:
// short (rel8) where available, otherwise the rel32 form. CALL is always
// rel32.
func jumpSize(mnem string, long bool) int {
if mnem == "CALL" {
return 5 // opcode + rel32
}
if !long {
return 2 // opcode + rel8
}
if mnem == "JMP" {
return 5 // E9 + rel32
}
return 6 // 0x0F 0x8x + rel32
}
// encodeInstr encodes one instruction, resolving jump targets against offsets
// (relative to pc, the instruction's own offset). A RET in a frame-pointer
// function is prefixed with the epilogue.
func encodeInstr(s *ast.Instr, pc int, offsets map[string]int, fi frameInfo) ([]byte, error) {
// function is prefixed with the epilogue. resolve, when non-nil, redirects a
// jump label through the jump-to-jump chain before the offset lookup.
func encodeInstr(s *ast.Instr, pc int, offsets map[string]int, fi frameInfo, long bool, resolve func(string) string) ([]byte, error) {
mnem := strings.ToUpper(s.Mnemonic.Text)
var prefix []byte
@@ -164,7 +265,7 @@ func encodeInstr(s *ast.Instr, pc int, offsets map[string]int, fi frameInfo) ([]
var code []byte
var err error
if isJumpMnemonic(mnem) {
code, err = encodeJump(s, mnem, pc+len(prefix), offsets)
code, err = encodeJump(s, mnem, pc+len(prefix), offsets, long, resolve)
} else {
code, err = encodeNormal(s, fi)
}
@@ -190,9 +291,9 @@ func encodeNormal(s *ast.Instr, fi frameInfo) ([]byte, error) {
return Encode(s.Mnemonic.Text, ops...)
}
// encodeJump encodes a JMP/CALL/Jcc with a rel32 offset resolved from the
// target label.
func encodeJump(s *ast.Instr, mnem string, pc int, offsets map[string]int) ([]byte, error) {
// encodeJump encodes a JMP/CALL/Jcc with a relative offset resolved from the
// target label, in the short (rel8) or long (rel32) form.
func encodeJump(s *ast.Instr, mnem string, pc int, offsets map[string]int, long bool, resolve func(string) string) ([]byte, error) {
if len(s.Operands) != 1 {
return nil, fmt.Errorf("jump expects 1 operand, got %d", len(s.Operands))
}
@@ -200,12 +301,25 @@ func encodeJump(s *ast.Instr, mnem string, pc int, offsets map[string]int) ([]by
if !ok {
return nil, fmt.Errorf("jump target must be a local label")
}
if resolve != nil && mnem != "CALL" {
name = resolve(name)
}
target, ok := offsets[name]
if !ok {
return nil, fmt.Errorf("undefined label %q", name)
}
rel := int64(target - (pc + jumpSize(mnem)))
rel := int64(target - (pc + jumpSize(mnem, long)))
if !long {
if !fits8(rel) {
return nil, fmt.Errorf("jump to %q does not fit the short form", name)
}
if mnem == "JMP" {
return []byte{0xEB, byte(int8(rel))}, nil
}
cc, _ := condCode(mnem)
return []byte{0x70 + byte(cc), byte(int8(rel))}, nil
}
switch mnem {
case "JMP":
return append([]byte{0xE9}, le32(rel)...), nil