// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: BSD-3-Clause package verify import ( "bytes" "math/rand" "testing" "unsafe" ) // decodeBlockGo is a minimal portable LZ4 block decoder used as the // differential-testing oracle. It mirrors the contract of // go-lz4's decodeBlockGo: (bytesWritten, code) where code is // 0 = ok, 1 = malformed, 2 = zero offset. func decodeBlockGo(src, dst []byte) (int, int) { if len(src) == 0 { return 0, 1 } si, di := 0, 0 for { if si >= len(src) { return 0, 1 // truncated: no token } token := int(src[si]) si++ // Literals. lLen := token >> 4 if lLen == 15 { for { if si >= len(src) { return 0, 1 } b := int(src[si]) si++ lLen += b if b != 255 { break } } } if si+lLen > len(src) { return 0, 1 // truncated literals } if di+lLen > len(dst) { return 0, 1 // destination overflow } copy(dst[di:di+lLen], src[si:si+lLen]) di += lLen si += lLen // End of block. if si >= len(src) { return di, 0 } // Match offset. if si+2 > len(src) { return 0, 1 } offset := int(src[si]) | int(src[si+1])<<8 si += 2 if offset == 0 { return 0, 2 } // Match length. mLen := token & 15 if mLen == 15 { for { if si >= len(src) { return 0, 1 } b := int(src[si]) si++ mLen += b if b != 255 { break } } } mLen += 4 // Copy match (overlapping-safe). if di-offset < 0 { return 0, 1 // offset reaches before dst start } if di+mLen > len(dst) { return 0, 1 // destination overflow } for i := 0; i < mLen; i++ { dst[di+i] = dst[di-offset+i] } di += mLen } } // genLZ4Block generates a random valid LZ4 block that decompresses into // approximately wantSize bytes. The block is always well-formed (ends with // a literals-only sequence). func genLZ4Block(rng *rand.Rand, wantSize int) []byte { var block []byte produced := 0 for produced < wantSize { remaining := wantSize - produced // Decide: emit a literals+match sequence or the final literals. if remaining <= 8 || rng.Intn(4) == 0 { // Final literals-only sequence. lLen := remaining if lLen > 60 { lLen = 1 + rng.Intn(60) } block = appendToken(block, lLen, 0) for i := 0; i < lLen; i++ { block = append(block, byte(rng.Intn(256))) } produced += lLen break } // Literals + match. lLen := rng.Intn(min(16, remaining)) if produced+lLen == 0 { lLen = 1 // must have at least 1 literal before the first match } mLenRaw := rng.Intn(12) // match length = mLenRaw + 4 mLen := mLenRaw + 4 if produced+mLen > remaining { mLen = remaining - produced if mLen < 4 { // Not enough room for a match; emit final literals. lLen = remaining block = appendToken(block, lLen, 0) for i := 0; i < lLen; i++ { block = append(block, byte(rng.Intn(256))) } break } mLenRaw = mLen - 4 } block = appendToken(block, lLen, mLenRaw) for i := 0; i < lLen; i++ { block = append(block, byte(rng.Intn(256))) } produced += lLen // Offset: must be <= produced (can't reference before start). maxOff := produced if maxOff > 65535 { maxOff = 65535 } offset := 1 + rng.Intn(maxOff) block = append(block, byte(offset), byte(offset>>8)) produced += mLen } return block } // appendToken appends a token (and extension bytes if needed) for the given // literal and match lengths. func appendToken(block []byte, lLen, mLenRaw int) []byte { lit4 := lLen if lit4 > 15 { lit4 = 15 } ml4 := mLenRaw if ml4 > 15 { ml4 = 15 } block = append(block, byte(lit4<<4|ml4)) // Literal extension bytes. rem := lLen - 15 for rem >= 255 { block = append(block, 255) rem -= 255 } if lLen >= 15 { block = append(block, byte(rem)) } // Match extension bytes. rem = mLenRaw - 15 for rem >= 255 { block = append(block, 255) rem -= 255 } if mLenRaw >= 15 { block = append(block, byte(rem)) } return block } func min(a, b int) int { if a < b { return a } return b } // TestDifferentialLZ4Fuzz drives the JIT-assembled decodeBlockAVX2 with // random valid LZ4 blocks and compares the output bit-for-bit against the // portable Go reference. func TestDifferentialLZ4Fuzz(t *testing.T) { k := loadLZ4Kernel(t) const iterations = 5000 rng := rand.New(rand.NewSource(42)) for i := 0; i < iterations; i++ { wantSize := 1 + rng.Intn(4096) src := genLZ4Block(rng, wantSize) dstSize := wantSize + 64 // generous destination // Go reference. goDst := make([]byte, dstSize) goN, goCode := decodeBlockGo(src, goDst) // JIT kernel. jitDst := make([]byte, dstSize) jitN, jitCode := callDecodeBlockAVX2(t, k, src, jitDst) if jitCode != goCode { t.Fatalf("iter %d: code mismatch: JIT=%d, Go=%d (src len=%d)", i, jitCode, goCode, len(src)) } if jitCode != 0 { continue // both agree it's malformed/zero-offset } if jitN != goN { t.Fatalf("iter %d: n mismatch: JIT=%d, Go=%d (src len=%d)", i, jitN, goN, len(src)) } if !bytes.Equal(jitDst[:jitN], goDst[:goN]) { t.Fatalf("iter %d: output mismatch (n=%d, src len=%d)", i, jitN, len(src)) } } } // TestDifferentialLZ4Hostile drives the kernel with random garbage to check // that error codes agree with the Go reference (no crashes, same classification). func TestDifferentialLZ4Hostile(t *testing.T) { k := loadLZ4Kernel(t) const iterations = 2000 rng := rand.New(rand.NewSource(99)) for i := 0; i < iterations; i++ { srcLen := rng.Intn(128) src := make([]byte, srcLen) rng.Read(src) dstSize := rng.Intn(512) dst := make([]byte, dstSize) // Go reference. goDst := make([]byte, dstSize) copy(goDst, dst) _, goCode := decodeBlockGo(src, goDst) // JIT kernel. jitDst := make([]byte, dstSize) copy(jitDst, dst) _, jitCode := callDecodeBlockAVX2(t, k, src, jitDst) if jitCode != goCode { t.Fatalf("iter %d: hostile code mismatch: JIT=%d, Go=%d (srcLen=%d, dstSize=%d)", i, jitCode, goCode, srcLen, dstSize) } } } // callDecodeBlockAVX2Raw is like callDecodeBlockAVX2 but accepts explicit // dst size (for hostile tests where dst may be smaller than the output). func callDecodeBlockAVX2Raw(t *testing.T, k *Kernel, src, dst []byte) (int, int) { t.Helper() 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 len(dst) > 0 { PutPtr(args, 24, unsafe.Pointer(&dst[0])) } PutUint64(args, 32, uint64(len(dst))) PutUint64(args, 40, uint64(cap(dst))) out, err := k.CallFunc("decodeBlockAVX2", args) if err != nil { t.Fatalf("CallFunc(decodeBlockAVX2): %v", err) } return int(GetUint64(out, 48)), int(GetUint64(out, 56)) }