fix(gasm): scaffold shared param names and two-sided seed sets

This commit is contained in:
2026-08-29 13:57:04 +02:00
parent 8bda4066e3
commit 6fb9629ab6
+82 -77
View File
@@ -9,7 +9,6 @@ import (
"go/parser" "go/parser"
"go/token" "go/token"
"os" "os"
"sort"
"strings" "strings"
gasmast "sourcedock.dev/petrbalvin/gasm-devkit/ast" gasmast "sourcedock.dev/petrbalvin/gasm-devkit/ast"
@@ -39,6 +38,11 @@ bodies, place the file in the kernel's package, and run it in CI.
return err return err
} }
rest := fs.Args() 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 { if len(rest) != 1 {
return fmt.Errorf("usage: gasm scaffold differential <file.s>") return fmt.Errorf("usage: gasm scaffold differential <file.s>")
} }
@@ -76,23 +80,30 @@ bodies, place the file in the kernel's package, and run it in CI.
fmt.Fprintf(&out, "func Test%sDifferential(t *testing.T) {\n", strings.ToUpper(name[:1])+name[1:]) 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, "\trng := rand.New(rand.NewSource(1))\n")
fmt.Fprintf(&out, "\tfor iter := 0; iter < 1000; iter++ {\n") fmt.Fprintf(&out, "\tfor iter := 0; iter < 1000; iter++ {\n")
written := map[string]bool{} // 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 { for _, p := range params {
genParamSeed(&out, p, written) a, b, slices := genParamSeed(&out, p, seen)
aArgs = append(aArgs, a)
bArgs = append(bArgs, b)
sliceNames = append(sliceNames, slices...)
} }
argList := make([]string, 0, len(params)) fmt.Fprintf(&out, "\t\tgot := %s(%s)\n", name, strings.Join(aArgs, ", "))
for _, p := range params { fmt.Fprintf(&out, "\t\twant := %sPortable(%s)\n", name, strings.Join(bArgs, ", "))
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\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\tt.Fatalf(\"iter %%d: kernel diverges from the portable spec\", iter)\n")
fmt.Fprintf(&out, "\t\t}\n\t}\n}\n\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(\"iter %%d: kernel mutated %%s differently\", iter, %q)\n", s)
fmt.Fprintf(&out, "\t\t}\n")
}
fmt.Fprintf(&out, "\t}\n}\n\n")
} }
if kernels == 0 { if kernels == 0 {
return fmt.Errorf("%s: no // func signatures found; add one doc comment per kernel", path) return fmt.Errorf("%s: no // func signatures found; add one doc comment per kernel", path)
@@ -144,7 +155,16 @@ func parseSig(doc string) ([]sigParam, []sigResult, bool) {
} }
var params []sigParam var params []sigParam
for _, field := range fd.Type.Params.List { for _, field := range fd.Type.Params.List {
params = append(params, sigParam{Names: identNames(field.Names), Type: exprString(field.Type)}) 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 var results []sigResult
if fd.Type.Results != nil { if fd.Type.Results != nil {
@@ -207,78 +227,63 @@ func resultDecl(results []sigResult) string {
return strings.Join(parts, ", ") return strings.Join(parts, ", ")
} }
// seedArg renders the argument expression referencing the seeded variables. // genParamSeed emits the seeding statements for one parameter and returns
func seedArg(p sigParam) string { // the kernel-side (A) and reference-side (B) argument expressions, plus the
if len(p.Names) == 0 { // names of any slice variables written in place (compared after the calls).
return typeSeedExpr(p.Type, "") 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, "[]")
} }
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 { switch {
case strings.HasPrefix(typ, "[]"): case isSlice:
if base == "" { v := uniqueName(seen, name)
base = "arg" 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)
return base + "Slice" fmt.Fprintf(out, "\t\tfor i := range %sA {\n", v)
case strings.HasPrefix(typ, "*"): fmt.Fprintf(out, "\t\t\tw%s := %s(rng.Intn(256))\n", v, goCast(elem))
if base == "" { fmt.Fprintf(out, "\t\t\t%sA[i] = w%s\n", v, v)
base = "arg" fmt.Fprintf(out, "\t\t\t%sB[i] = w%s\n", v, v)
} fmt.Fprintf(out, "\t\t}\n")
return "&" + base + "Elem" 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: default:
if base == "" { v := uniqueName(seen, name)
base = "arg" 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 base + "Scalar" return v + "A", v + "B", nil
} }
} }
// genParamSeed emits the seeding statements for one parameter. // uniqueName de-duplicates seeded variable names when one kernel takes two
func genParamSeed(out *strings.Builder, p sigParam, written map[string]bool) { // parameters of the same name (impossible in Go) or a name repeats across
names := sortedUnique(p.Names) // kernels in one file.
for _, n := range names { func uniqueName(seen map[string]bool, base string) string {
switch { if base == "" {
case strings.HasPrefix(p.Type, "[]"): base = "arg"
if written[n+"Slice"] { }
continue if !seen[base] {
} seen[base] = true
written[n+"Slice"] = true return base
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, "[]"))) for i := 2; ; i++ {
case strings.HasPrefix(p.Type, "*"): cand := fmt.Sprintf("%s%d", base, i)
if written[n+"Elem"] { if !seen[cand] {
continue seen[cand] = true
} return cand
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 { // goCast returns the conversion turning rng.Intn into the element type.
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 { func goCast(elem string) string {
switch elem { switch elem {
case "byte", "uint8": case "byte", "uint8":