330 lines
9.1 KiB
Go
330 lines
9.1 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"
|
|
"sort"
|
|
"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()
|
|
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")
|
|
|
|
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 iter := 0; iter < 1000; iter++ {\n")
|
|
written := map[string]bool{}
|
|
for _, p := range params {
|
|
genParamSeed(&out, p, written)
|
|
}
|
|
argList := make([]string, 0, len(params))
|
|
for _, p := range params {
|
|
argList = append(argList, seedArg(p))
|
|
}
|
|
fmt.Fprintf(&out, "\t\tgot := %s(%s)\n", name, strings.Join(argList, ", "))
|
|
wantArgs := make([]string, 0, len(params))
|
|
for _, p := range params {
|
|
wantArgs = append(wantArgs, seedArg(p))
|
|
}
|
|
fmt.Fprintf(&out, "\t\twant := %sPortable(%s)\n", name, strings.Join(wantArgs, ", "))
|
|
fmt.Fprintf(&out, "\t\tif !bytes.Equal(outputBytes(got), outputBytes(want)) {\n")
|
|
fmt.Fprintf(&out, "\t\t\tt.Fatalf(\"iter %%d: kernel diverges from the portable spec\", iter)\n")
|
|
fmt.Fprintf(&out, "\t\t}\n\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.Split(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 {
|
|
params = append(params, sigParam{Names: identNames(field.Names), Type: exprString(field.Type)})
|
|
}
|
|
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, ", ")
|
|
}
|
|
|
|
// seedArg renders the argument expression referencing the seeded variables.
|
|
func seedArg(p sigParam) string {
|
|
if len(p.Names) == 0 {
|
|
return typeSeedExpr(p.Type, "")
|
|
}
|
|
return typeSeedExpr(p.Type, p.Names[0])
|
|
}
|
|
|
|
// typeSeedExpr returns the expression passed for one parameter, given the
|
|
// seeded variable base name ("" for nameless parameters).
|
|
func typeSeedExpr(typ, base string) string {
|
|
switch {
|
|
case strings.HasPrefix(typ, "[]"):
|
|
if base == "" {
|
|
base = "arg"
|
|
}
|
|
return base + "Slice"
|
|
case strings.HasPrefix(typ, "*"):
|
|
if base == "" {
|
|
base = "arg"
|
|
}
|
|
return "&" + base + "Elem"
|
|
default:
|
|
if base == "" {
|
|
base = "arg"
|
|
}
|
|
return base + "Scalar"
|
|
}
|
|
}
|
|
|
|
// genParamSeed emits the seeding statements for one parameter.
|
|
func genParamSeed(out *strings.Builder, p sigParam, written map[string]bool) {
|
|
names := sortedUnique(p.Names)
|
|
for _, n := range names {
|
|
switch {
|
|
case strings.HasPrefix(p.Type, "[]"):
|
|
if written[n+"Slice"] {
|
|
continue
|
|
}
|
|
written[n+"Slice"] = true
|
|
fmt.Fprintf(out, "\t\t%sSlice := make([]%s, 1+rng.Intn(512))\n", n, strings.TrimPrefix(p.Type, "[]"))
|
|
fmt.Fprintf(out, "\t\tfor i := range %sSlice {\n\t\t\t%sSlice[i] = %s(rng.Intn(256))\n\t\t}\n", n, strings.TrimPrefix(p.Type, "[]"), goCast(strings.TrimPrefix(p.Type, "[]")))
|
|
case strings.HasPrefix(p.Type, "*"):
|
|
if written[n+"Elem"] {
|
|
continue
|
|
}
|
|
written[n+"Elem"] = true
|
|
fmt.Fprintf(out, "\t\tvar %sElem %s\n\t\t%sElem = %s(rng.Intn(256))\n", n, strings.TrimPrefix(p.Type, "*"), n, goCast(strings.TrimPrefix(p.Type, "*")))
|
|
default:
|
|
if written[n+"Scalar"] {
|
|
continue
|
|
}
|
|
written[n+"Scalar"] = true
|
|
fmt.Fprintf(out, "\t\t%sScalar := rng.Intn(512)\n", n)
|
|
}
|
|
}
|
|
}
|
|
|
|
func sortedUnique(names []string) []string {
|
|
seen := map[string]bool{}
|
|
var out []string
|
|
for _, n := range names {
|
|
if !seen[n] {
|
|
seen[n] = true
|
|
out = append(out, n)
|
|
}
|
|
}
|
|
sort.Strings(out)
|
|
return out
|
|
}
|
|
|
|
// goCast returns the conversion turning rng.Intn(256) 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"
|
|
}
|
|
}
|
|
|
|
// outputBytes narrows a returned slice to bytes for the comparison; scalar
|
|
// results are compared through the same helper via a reflect-free trick the
|
|
// author may need to adjust for non-slice returns.
|
|
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*4/2] = byte(x)
|
|
b[i*4/2+1] = byte(x >> 8)
|
|
}
|
|
return b
|
|
default:
|
|
return nil
|
|
}
|
|
}
|