// Copyright (c) 2026 Petr Balvín (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 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 ", ` Print a differential test skeleton for every // func signature in FILE. The test seeds random states, drives the kernel and a portable reference (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 ") } 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 } }