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/token"
"os"
"sort"
"strings"
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
}
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>")
}
@@ -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, "\trng := rand.New(rand.NewSource(1))\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 {
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))
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\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(\"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 {
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
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
if fd.Type.Results != nil {
@@ -207,78 +227,63 @@ func resultDecl(results []sigResult) string {
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, "")
// 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, "[]")
}
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"
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:
if base == "" {
base = "arg"
}
return base + "Scalar"
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
}
}
// 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)
// 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
}
}
}
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.
// goCast returns the conversion turning rng.Intn into the element type.
func goCast(elem string) string {
switch elem {
case "byte", "uint8":