145 lines
3.6 KiB
Go
145 lines
3.6 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: BSD-3-Clause
|
||
|
|
|
||
|
|
package verify
|
||
|
|
|
||
|
|
import (
|
||
|
|
"bytes"
|
||
|
|
"math/rand"
|
||
|
|
"os"
|
||
|
|
"testing"
|
||
|
|
"unsafe"
|
||
|
|
)
|
||
|
|
|
||
|
|
const lz4AVX512Path = "../../go-libraries/go-lz4/avx512_amd64.s"
|
||
|
|
|
||
|
|
func loadLZ4AVX512Kernel(t *testing.T) *Kernel {
|
||
|
|
t.Helper()
|
||
|
|
if _, err := os.Stat(lz4AVX512Path); err != nil {
|
||
|
|
t.Skipf("sibling kernel not available: %v", err)
|
||
|
|
}
|
||
|
|
k, err := Load(lz4AVX512Path)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Load(%s): %v", lz4AVX512Path, err)
|
||
|
|
}
|
||
|
|
t.Cleanup(k.Close)
|
||
|
|
return k
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAVX512DecodeKnownAnswers(t *testing.T) {
|
||
|
|
k := loadLZ4AVX512Kernel(t)
|
||
|
|
|
||
|
|
tests := []struct {
|
||
|
|
name string
|
||
|
|
src []byte
|
||
|
|
wantN int
|
||
|
|
wantCode int
|
||
|
|
}{
|
||
|
|
{"literals_only", []byte{0x50, 'H', 'e', 'l', 'l', 'o'}, 5, 0},
|
||
|
|
{"literals_and_match", []byte{0x54, 'A', 'A', 'A', 'A', 'A', 0x05, 0x00, 0x30, 'B', 'B', 'B'}, 16, 0},
|
||
|
|
{"overlapping", []byte{0x14, 'X', 0x01, 0x00, 0x10, 'Y'}, 10, 0},
|
||
|
|
{"malformed", []byte{0x50, 'H', 'e'}, 0, 1},
|
||
|
|
{"zero_offset", []byte{0x14, 'X', 0x00, 0x00}, 0, 2},
|
||
|
|
}
|
||
|
|
for _, tt := range tests {
|
||
|
|
t.Run(tt.name, func(t *testing.T) {
|
||
|
|
dst := make([]byte, 64)
|
||
|
|
args := make([]byte, 64)
|
||
|
|
PutPtr(args, 0, unsafe.Pointer(&tt.src[0]))
|
||
|
|
PutUint64(args, 8, uint64(len(tt.src)))
|
||
|
|
PutUint64(args, 16, uint64(cap(tt.src)))
|
||
|
|
PutPtr(args, 24, unsafe.Pointer(&dst[0]))
|
||
|
|
PutUint64(args, 32, uint64(len(dst)))
|
||
|
|
PutUint64(args, 40, uint64(cap(dst)))
|
||
|
|
|
||
|
|
out, err := k.CallFunc("decodeBlockAVX512", args)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("CallFunc: %v", err)
|
||
|
|
}
|
||
|
|
n := int(GetUint64(out, 48))
|
||
|
|
code := int(GetUint64(out, 56))
|
||
|
|
if n != tt.wantN || code != tt.wantCode {
|
||
|
|
t.Errorf("got (n=%d, code=%d), want (n=%d, code=%d)", n, code, tt.wantN, tt.wantCode)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAVX512DifferentialFuzz(t *testing.T) {
|
||
|
|
k := loadLZ4AVX512Kernel(t)
|
||
|
|
rng := rand.New(rand.NewSource(77))
|
||
|
|
|
||
|
|
for i := 0; i < 3000; i++ {
|
||
|
|
wantSize := 1 + rng.Intn(4096)
|
||
|
|
src := genLZ4Block(rng, wantSize)
|
||
|
|
dstSize := wantSize + 64
|
||
|
|
|
||
|
|
goDst := make([]byte, dstSize)
|
||
|
|
goN, goCode := decodeBlockGo(src, goDst)
|
||
|
|
|
||
|
|
jitDst := make([]byte, dstSize)
|
||
|
|
args := make([]byte, 64)
|
||
|
|
if len(src) > 0 {
|
||
|
|
PutPtr(args, 0, unsafe.Pointer(&src[0]))
|
||
|
|
}
|
||
|
|
PutUint64(args, 8, uint64(len(src)))
|
||
|
|
PutUint64(args, 16, uint64(cap(src)))
|
||
|
|
if dstSize > 0 {
|
||
|
|
PutPtr(args, 24, unsafe.Pointer(&jitDst[0]))
|
||
|
|
}
|
||
|
|
PutUint64(args, 32, uint64(dstSize))
|
||
|
|
PutUint64(args, 40, uint64(cap(jitDst)))
|
||
|
|
|
||
|
|
out, err := k.CallFunc("decodeBlockAVX512", args)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("iter %d: %v", i, err)
|
||
|
|
}
|
||
|
|
jitN := int(GetUint64(out, 48))
|
||
|
|
jitCode := int(GetUint64(out, 56))
|
||
|
|
|
||
|
|
if jitCode != goCode {
|
||
|
|
t.Fatalf("iter %d: code mismatch: JIT=%d Go=%d", i, jitCode, goCode)
|
||
|
|
}
|
||
|
|
if jitCode != 0 {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
if jitN != goN {
|
||
|
|
t.Fatalf("iter %d: n mismatch: JIT=%d Go=%d", i, jitN, goN)
|
||
|
|
}
|
||
|
|
if !bytes.Equal(jitDst[:jitN], goDst[:goN]) {
|
||
|
|
t.Fatalf("iter %d: output mismatch (n=%d)", i, jitN)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAVX512WideCopy(t *testing.T) {
|
||
|
|
k := loadLZ4AVX512Kernel(t)
|
||
|
|
|
||
|
|
sizes := []int{0, 1, 31, 32, 63, 64, 65, 127, 128, 256, 1024}
|
||
|
|
for _, n := range sizes {
|
||
|
|
src := make([]byte, n)
|
||
|
|
for i := range src {
|
||
|
|
src[i] = byte(i*11 + 3)
|
||
|
|
}
|
||
|
|
dst := make([]byte, n)
|
||
|
|
|
||
|
|
args := make([]byte, 48)
|
||
|
|
if n > 0 {
|
||
|
|
PutPtr(args, 0, unsafe.Pointer(&dst[0]))
|
||
|
|
PutPtr(args, 24, unsafe.Pointer(&src[0]))
|
||
|
|
}
|
||
|
|
PutUint64(args, 8, uint64(n))
|
||
|
|
PutUint64(args, 16, uint64(n))
|
||
|
|
PutUint64(args, 32, uint64(n))
|
||
|
|
PutUint64(args, 40, uint64(n))
|
||
|
|
|
||
|
|
_, err := k.CallFunc("wideCopyAVX512", args)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("wideCopyAVX512(n=%d): %v", n, err)
|
||
|
|
}
|
||
|
|
if !bytes.Equal(dst, src) {
|
||
|
|
t.Errorf("wideCopyAVX512(n=%d): mismatch", n)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|