346 lines
10 KiB
Go
346 lines
10 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: BSD-3-Clause
|
|
|
|
package main
|
|
|
|
import (
|
|
"fmt"
|
|
"go/ast"
|
|
"go/parser"
|
|
"go/token"
|
|
"os"
|
|
"strings"
|
|
|
|
gasmast "sourcedock.dev/petrbalvin/gasm-devkit/ast"
|
|
gasmparser "sourcedock.dev/petrbalvin/gasm-devkit/parser"
|
|
)
|
|
|
|
// cmdScaffold generates a differential test skeleton for every kernel in a
|
|
// file: a Go test that seeds random states, drives both the assembly kernel
|
|
// and a caller-provided portable reference, and compares the outputs
|
|
// byte-for-byte. The lesson this encodes: a pipeline-level fuzz cannot see
|
|
// an unwired kernel — only a direct-call differential against the portable
|
|
// specification can, so every kernel ships with one.
|
|
//
|
|
// The generated file follows two conventions the caller fills in:
|
|
// - the assembly symbols resolve because the test lives in the kernel's
|
|
// own package (the //go:noescape declarations reference them);
|
|
// - each kernel gets a <name>Portable Go function the author implements as
|
|
// the specification, and the test fails on the first divergent byte.
|
|
func cmdScaffold(args []string) error {
|
|
fs := newCommand("scaffold", "gasm scaffold differential <file.s>", `
|
|
Print a differential test skeleton for every // func signature in FILE.
|
|
The test seeds random states, drives the kernel and a portable reference
|
|
(<name>Portable), and compares outputs byte-for-byte. Write the reference
|
|
bodies, place the file in the kernel's package, and run it in CI.
|
|
`)
|
|
if err := fs.Parse(args); err != nil {
|
|
return err
|
|
}
|
|
rest := fs.Args()
|
|
// The first positional word is the scaffold style; "differential" is the
|
|
// only one today.
|
|
if len(rest) > 0 && rest[0] == "differential" {
|
|
rest = rest[1:]
|
|
}
|
|
if len(rest) != 1 {
|
|
return fmt.Errorf("usage: gasm scaffold differential <file.s>")
|
|
}
|
|
path := rest[0]
|
|
src, err := os.ReadFile(path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
f, errs := gasmparser.Parse(path, string(src))
|
|
if len(errs) > 0 {
|
|
return fmt.Errorf("parse: %v", errs[0])
|
|
}
|
|
|
|
var out strings.Builder
|
|
out.WriteString(headerComment)
|
|
out.WriteString("package " + packageName + "\n\n")
|
|
out.WriteString("import (\n\t\"bytes\"\n\t\"math/rand\"\n\t\"testing\"\n)\n\n")
|
|
out.WriteString(generatedHelpers)
|
|
|
|
kernels := 0
|
|
for _, d := range f.Decls {
|
|
txt, ok := d.(*gasmast.Text)
|
|
if !ok {
|
|
continue
|
|
}
|
|
params, results, ok := parseSig(txt.Doc)
|
|
if !ok || len(params) == 0 {
|
|
continue
|
|
}
|
|
kernels++
|
|
name := txt.Name.Name
|
|
fmt.Fprintf(&out, "// %sPortable is the specification %s is pinned against:\n", name, name)
|
|
fmt.Fprintf(&out, "// fill in a straightforward implementation of the same contract.\n")
|
|
fmt.Fprintf(&out, "func %sPortable(%s) (%s) {\n\tpanic(\"implement the portable specification\")\n}\n\n", name, paramDecl(params), resultDecl(results))
|
|
|
|
fmt.Fprintf(&out, "func Test%sDifferential(t *testing.T) {\n", strings.ToUpper(name[:1])+name[1:])
|
|
fmt.Fprintf(&out, "\trng := rand.New(rand.NewSource(1))\n")
|
|
fmt.Fprintf(&out, "\tfor range 1000 {\n")
|
|
// Seed two independent argument sets per iteration: the kernel runs
|
|
// on set A, the portable reference on set B, so in-place writes
|
|
// through pointer/slice arguments cannot contaminate the other side.
|
|
var sliceNames []string
|
|
seen := map[string]bool{}
|
|
aArgs := make([]string, 0, len(params))
|
|
bArgs := make([]string, 0, len(params))
|
|
for _, p := range params {
|
|
a, b, slices := genParamSeed(&out, p, seen)
|
|
aArgs = append(aArgs, a)
|
|
bArgs = append(bArgs, b)
|
|
sliceNames = append(sliceNames, slices...)
|
|
}
|
|
fmt.Fprintf(&out, "\t\tgot := %s(%s)\n", name, strings.Join(aArgs, ", "))
|
|
fmt.Fprintf(&out, "\t\twant := %sPortable(%s)\n", name, strings.Join(bArgs, ", "))
|
|
fmt.Fprintf(&out, "\t\tif !bytes.Equal(outputBytes(got), outputBytes(want)) {\n")
|
|
fmt.Fprintf(&out, "\t\t\tt.Fatalf(\"kernel diverges from the portable spec (seed 1, deterministic)\")\n")
|
|
fmt.Fprintf(&out, "\t\t}\n")
|
|
for _, s := range sliceNames {
|
|
fmt.Fprintf(&out, "\t\tif !bytes.Equal(outputBytes(%sA), outputBytes(%sB)) {\n", s, s)
|
|
fmt.Fprintf(&out, "\t\t\tt.Fatalf(\"kernel mutated %%q differently (seed 1, deterministic)\", %q)\n", s)
|
|
fmt.Fprintf(&out, "\t\t}\n")
|
|
}
|
|
fmt.Fprintf(&out, "\t}\n}\n\n")
|
|
}
|
|
if kernels == 0 {
|
|
return fmt.Errorf("%s: no // func signatures found; add one doc comment per kernel", path)
|
|
}
|
|
os.Stdout.WriteString(out.String())
|
|
return nil
|
|
}
|
|
|
|
const packageName = "yourpkg"
|
|
|
|
const headerComment = `// Code generated by gasm scaffold differential; EDIT THE PANICS.
|
|
// Each Test*Differential drives the assembly kernel and its portable
|
|
// reference over the same random states and compares the outputs.
|
|
// Place this file in the kernel's own package so the symbols resolve.
|
|
|
|
`
|
|
|
|
// sigParam is one parsed // func parameter.
|
|
type sigParam struct {
|
|
Names []string
|
|
Type string
|
|
}
|
|
|
|
type sigResult struct {
|
|
Names []string
|
|
Type string
|
|
}
|
|
|
|
// parseSig parses the // func signature of a doc comment.
|
|
func parseSig(doc string) ([]sigParam, []sigResult, bool) {
|
|
var line string
|
|
for l := range strings.SplitSeq(doc, "\n") {
|
|
if t := strings.TrimSpace(l); strings.HasPrefix(t, "func ") {
|
|
line = t
|
|
break
|
|
}
|
|
}
|
|
if line == "" {
|
|
return nil, nil, false
|
|
}
|
|
fset := token.NewFileSet()
|
|
f, err := parser.ParseFile(fset, "sig.go", "package p\n"+line+" {}\n", 0)
|
|
if err != nil {
|
|
return nil, nil, false
|
|
}
|
|
fd, ok := f.Decls[0].(*ast.FuncDecl)
|
|
if !ok || fd.Type == nil {
|
|
return nil, nil, false
|
|
}
|
|
var params []sigParam
|
|
for _, field := range fd.Type.Params.List {
|
|
typ := exprString(field.Type)
|
|
if len(field.Names) == 0 {
|
|
params = append(params, sigParam{Names: []string{""}, Type: typ})
|
|
continue
|
|
}
|
|
// Shared names (`L, result *byte`) expand to one entry per name:
|
|
// every name is a separate argument at the call site.
|
|
for _, n := range field.Names {
|
|
params = append(params, sigParam{Names: []string{n.Name}, Type: typ})
|
|
}
|
|
}
|
|
var results []sigResult
|
|
if fd.Type.Results != nil {
|
|
for _, field := range fd.Type.Results.List {
|
|
results = append(results, sigResult{Names: identNames(field.Names), Type: exprString(field.Type)})
|
|
}
|
|
}
|
|
return params, results, true
|
|
}
|
|
|
|
func identNames(idents []*ast.Ident) []string {
|
|
var out []string
|
|
for _, id := range idents {
|
|
out = append(out, id.Name)
|
|
}
|
|
return out
|
|
}
|
|
|
|
func exprString(e ast.Expr) string {
|
|
switch t := e.(type) {
|
|
case *ast.Ident:
|
|
return t.Name
|
|
case *ast.StarExpr:
|
|
return "*" + exprString(t.X)
|
|
case *ast.SelectorExpr:
|
|
return exprString(t.X) + "." + t.Sel.Name
|
|
case *ast.ArrayType:
|
|
if t.Len == nil {
|
|
return "[]" + exprString(t.Elt)
|
|
}
|
|
return "[N]" + exprString(t.Elt)
|
|
}
|
|
return "interface{}"
|
|
}
|
|
|
|
// paramDecl renders a parameter list for the portable reference signature.
|
|
func paramDecl(params []sigParam) string {
|
|
var parts []string
|
|
for _, p := range params {
|
|
if len(p.Names) == 0 {
|
|
parts = append(parts, p.Type)
|
|
continue
|
|
}
|
|
for _, n := range p.Names {
|
|
parts = append(parts, n+" "+p.Type)
|
|
}
|
|
}
|
|
return strings.Join(parts, ", ")
|
|
}
|
|
|
|
// resultDecl renders a result list; unnamed results keep bare types.
|
|
func resultDecl(results []sigResult) string {
|
|
if len(results) == 0 {
|
|
return ""
|
|
}
|
|
var parts []string
|
|
for _, r := range results {
|
|
parts = append(parts, r.Type)
|
|
}
|
|
return strings.Join(parts, ", ")
|
|
}
|
|
|
|
// genParamSeed emits the seeding statements for one parameter and returns
|
|
// the kernel-side (A) and reference-side (B) argument expressions, plus the
|
|
// names of any slice variables written in place (compared after the calls).
|
|
func genParamSeed(out *strings.Builder, p sigParam, seen map[string]bool) (aArg, bArg string, slices []string) {
|
|
name := p.Names[0]
|
|
elem := strings.TrimPrefix(p.Type, "*")
|
|
isSlice := strings.HasPrefix(p.Type, "[]")
|
|
if isSlice {
|
|
elem = strings.TrimPrefix(p.Type, "[]")
|
|
}
|
|
switch {
|
|
case isSlice:
|
|
v := uniqueName(seen, name)
|
|
fmt.Fprintf(out, "\t\t%sA := make([]%s, 1+rng.Intn(512))\n", v, elem)
|
|
fmt.Fprintf(out, "\t\t%sB := make([]%s, len(%sA))\n", v, elem, v)
|
|
fmt.Fprintf(out, "\t\tfor i := range %sA {\n", v)
|
|
fmt.Fprintf(out, "\t\t\tw%s := %s(rng.Intn(256))\n", v, goCast(elem))
|
|
fmt.Fprintf(out, "\t\t\t%sA[i] = w%s\n", v, v)
|
|
fmt.Fprintf(out, "\t\t\t%sB[i] = w%s\n", v, v)
|
|
fmt.Fprintf(out, "\t\t}\n")
|
|
return v, v, []string{v}
|
|
case strings.HasPrefix(p.Type, "*"):
|
|
v := uniqueName(seen, name)
|
|
fmt.Fprintf(out, "\t\tvar %sA, %sB %s\n", v, v, elem)
|
|
fmt.Fprintf(out, "\t\tw%s := %s(rng.Intn(256))\n", v, goCast(elem))
|
|
fmt.Fprintf(out, "\t\t%sA = w%s\n", v, v)
|
|
fmt.Fprintf(out, "\t\t%sB = w%s\n", v, v)
|
|
return "&" + v + "A", "&" + v + "B", nil
|
|
default:
|
|
v := uniqueName(seen, name)
|
|
fmt.Fprintf(out, "\t\tw%s := %s(rng.Intn(512))\n", v, goCast(""))
|
|
fmt.Fprintf(out, "\t\tvar %sA, %sB %s = w%s, w%s\n", v, v, p.Type, v, v)
|
|
return v + "A", v + "B", nil
|
|
}
|
|
}
|
|
|
|
// uniqueName de-duplicates seeded variable names when one kernel takes two
|
|
// parameters of the same name (impossible in Go) or a name repeats across
|
|
// kernels in one file.
|
|
func uniqueName(seen map[string]bool, base string) string {
|
|
if base == "" {
|
|
base = "arg"
|
|
}
|
|
if !seen[base] {
|
|
seen[base] = true
|
|
return base
|
|
}
|
|
for i := 2; ; i++ {
|
|
cand := fmt.Sprintf("%s%d", base, i)
|
|
if !seen[cand] {
|
|
seen[cand] = true
|
|
return cand
|
|
}
|
|
}
|
|
}
|
|
|
|
// goCast returns the conversion turning rng.Intn into the element type.
|
|
func goCast(elem string) string {
|
|
switch elem {
|
|
case "byte", "uint8":
|
|
return "byte"
|
|
case "int8":
|
|
return "int8"
|
|
case "uint16":
|
|
return "uint16"
|
|
case "int16":
|
|
return "int16"
|
|
case "uint32":
|
|
return "uint32"
|
|
case "int32":
|
|
return "int32"
|
|
case "uint64":
|
|
return "uint64"
|
|
default:
|
|
return "int"
|
|
}
|
|
}
|
|
|
|
// generatedHelpers is emitted into every generated test file: outputBytes
|
|
// narrows returned slices and scalars to a byte form for the comparison.
|
|
// It lives in the template, not in this binary, because only the generated
|
|
// file ever calls it.
|
|
const generatedHelpers = `// outputBytes narrows a returned slice or scalar to bytes for the
|
|
// comparison; extend the switch when a kernel returns a wider type.
|
|
func outputBytes(v any) []byte {
|
|
switch t := v.(type) {
|
|
case []byte:
|
|
return t
|
|
case []int32:
|
|
b := make([]byte, 4*len(t))
|
|
for i, x := range t {
|
|
b[i*4] = byte(x)
|
|
b[i*4+1] = byte(x >> 8)
|
|
b[i*4+2] = byte(x >> 16)
|
|
b[i*4+3] = byte(x >> 24)
|
|
}
|
|
return b
|
|
case []uint16:
|
|
b := make([]byte, 2*len(t))
|
|
for i, x := range t {
|
|
b[i*2] = byte(x)
|
|
b[i*2+1] = byte(x >> 8)
|
|
}
|
|
return b
|
|
case int:
|
|
b := make([]byte, 8)
|
|
for i := range 8 {
|
|
b[i] = byte(uint64(t) >> (8 * i))
|
|
}
|
|
return b
|
|
default:
|
|
return nil
|
|
}
|
|
}
|
|
`
|