diff --git a/cmd/gasm/scaffold.go b/cmd/gasm/scaffold.go index 9460f97..232d1c3 100644 --- a/cmd/gasm/scaffold.go +++ b/cmd/gasm/scaffold.go @@ -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 ") } @@ -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":