diff --git a/internal/cuetest/cuetest.go b/internal/cuetest/cuetest.go index f52faa609..52242aa63 100644 --- a/internal/cuetest/cuetest.go +++ b/internal/cuetest/cuetest.go @@ -21,6 +21,8 @@ import ( "os" "regexp" "testing" + + "cuelang.org/go/internal/tdtest" ) const ( @@ -95,6 +97,18 @@ func Condition(cond string) (bool, error) { return false, fmt.Errorf("unknown condition %v", cond) } +// T is an alias to tdtest.T +type T = tdtest.T + +func init() { + tdtest.UpdateTests = UpdateGoldenFiles +} + +// Run creates a new table-driven test using the CUE testing defaults. +func Run[TC any](t *testing.T, table []TC, fn func(t *T, tc *TC)) { + tdtest.Run(t, table, fn) +} + // IssueSkip causes the test t to be skipped unless the issue identified // by s is deemed to be a non-issue by CUE_NON_ISSUES. func IssueSkip(t *testing.T, s string) { diff --git a/internal/tdtest/tdtest.go b/internal/tdtest/tdtest.go new file mode 100644 index 000000000..f2eaca6cc --- /dev/null +++ b/internal/tdtest/tdtest.go @@ -0,0 +1,175 @@ +// Copyright 2023 CUE Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Package tdtest provides support for table-driven testing. +// +// Features include automatically updating of test values, automatic error +// message generation, and singling out single tests to run. +// +// Auto updating fields is only supported for fields that are scalar types: +// string, bool, int*, and uint*. If the field is a string, the "actual" value +// may be any Go value that can meaningfully be printed with fmt.Sprint. +package tdtest + +import ( + "fmt" + "go/token" + "reflect" + "runtime" + "strings" + "testing" +) + +// TODO: +// - make this a public package at some point. +// - add tests. Maybe adding Examples is sufficient. +// - use text-based modification, instead of astutil. The latter is too brittle. +// - allow updating position-based, instead of named, fields. +// - implement skip, maybe match +// - make name field explicit, i.e. Name("name"), field tag, or tdtest.Name type. +// - allow "skip" field. Again either SkipName("skip"), tag, or Skip type. +// - allow for tdtest:"noupdate" field tag. +// - should we derive names from field names? This would require always +// loading the packages data upon error. Could be an option to disable, or +// implicitly it would only be loaded if there is an error without message. +// - Option: allow ignore field that lists a set of fields to not be tested +// for that particular test case: ignore: tdtest.Ignore("want1", "want2") +// + +// UpdateTests defines whether tests should be updated by default. +// This can be overridden on an individual basis using T.Update. +var UpdateTests = false + +// set is the set of tests to run. +type set[TC any] struct { + t *testing.T + + table []TC + toRun []int + + updateEnabled bool + file string + info *info +} + +// Run runs the given function for each (selected) element in the table. +func Run[TC any](t *testing.T, table []TC, fn func(t *T, tc *TC)) { + s := &set[TC]{ + t: t, + table: table, + updateEnabled: UpdateTests, + } + for i := range s.table { + name := fmt.Sprint(i) + + x := reflect.ValueOf(s.table[i]).FieldByName("name") + if x.Kind() == reflect.String { + name += "/" + x.String() + } + + s.t.Run(name, func(t *testing.T) { + tt := &T{ + T: t, + iter: i, + infoSrc: s, + updateEnabled: s.updateEnabled, + } + fn(tt, &s.table[i]) + }) + } + if s.info != nil && s.info.needsUpdate { + s.update() + } +} + +// T is a single test case representing an element in a table. +// It embeds *testing.T, so all functions of testing.T are available. +type T struct { + *testing.T + + infoSrc interface{ getInfo(file string) *info } + iter int // position in the table of the current subtest. + + updateEnabled bool +} + +func (t *T) info(file string) *info { + return t.infoSrc.getInfo(file) +} + +func (t *T) getCallInfo() (*info, *callInfo) { + _, file, line, ok := runtime.Caller(2) + if !ok { + t.Fatalf("could not update file for test %s", t.Name()) + } + info := t.info(file) + return info, info.calls[token.Position{Filename: file, Line: line}] +} + +// Equal compares two fields. +// +// For auto updating to work, field must reference a field in the test case +// directly. +func (t *T) Equal(actual, field any, msgAndArgs ...any) { + t.Helper() + + switch { + case field == actual: + case t.updateEnabled: + info, ci := t.getCallInfo() + t.updateField(info, ci, actual) + case len(msgAndArgs) == 0: + t.Errorf("unexpected value:\ngot: %v;\nwant: %v", actual, field) + default: + format := msgAndArgs[0].(string) + ":\ngot: %v;\nwant: %v" + args := append(msgAndArgs[1:], actual, field) + t.Errorf(format, args...) + } +} + +// Update specifies whether to update the Go structs in case of discrepancies. +// It overrides the default setting. +func (t *T) Update(enable bool) { + t.updateEnabled = enable +} + +// Select species which tests to run. The test may be an int, in which case +// it selects the table entry to run, or a string, which is matched against +// the last path of the test. An empty list runs all tests. +func (t *T) Select(tests ...any) { + if len(tests) == 0 { + return + } + + t.Helper() + + name := t.Name() + parts := strings.Split(name, "/") + + for _, n := range tests { + switch n := n.(type) { + case int: + if n == t.iter { + return + } + case string: + if n == parts[len(parts)-1] { + return + } + default: + panic("unexpected type passed to Select") + } + } + t.Skip("not selected") +} diff --git a/internal/tdtest/tdtest_test.go b/internal/tdtest/tdtest_test.go new file mode 100644 index 000000000..0e6c38b5f --- /dev/null +++ b/internal/tdtest/tdtest_test.go @@ -0,0 +1,42 @@ +// Copyright 2023 CUE Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package tdtest_test + +import ( + "testing" + + "cuelang.org/go/internal/tdtest" +) + +// TODO: write a proper test + +// NOTE: for debugging purposes. Do not remove. +func TestX(t *testing.T) { + t.Skip() + + type testCase struct { + name string + want string + } + _, cases := 1, []testCase{{ + name: "foo", + want: `foo`, + }} + + tdtest.Run(t, cases, func(t *tdtest.T, tc *testCase) { + t.Update(true) + t.Equal("actual", tc.want) + }) +} diff --git a/internal/tdtest/update.go b/internal/tdtest/update.go new file mode 100644 index 000000000..e3f9b96d7 --- /dev/null +++ b/internal/tdtest/update.go @@ -0,0 +1,402 @@ +// Copyright 2023 CUE Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package tdtest + +import ( + "fmt" + "go/ast" + "go/format" + "go/token" + "go/types" + "os" + "reflect" + "strconv" + "strings" + "sync" + "testing" + + "golang.org/x/tools/go/ast/astutil" + "golang.org/x/tools/go/packages" +) + +// info contains information needed to update files. +type info struct { + t *testing.T + + tcType reflect.Type + + needsUpdate bool // an updateable field has changed + + table *ast.CompositeLit // the table that is the source of the tests + + testPkg *packages.Package + + calls map[token.Position]*callInfo + patches map[ast.Node]ast.Expr +} + +type callInfo struct { + ast *ast.CallExpr + funcName string + fieldName string +} + +var ( + once sync.Once + pkgs []*packages.Package + pkgsErr error +) + +func initPackages() ([]*packages.Package, error) { + once.Do(func() { + cfg := &packages.Config{ + Mode: packages.NeedFiles | + packages.NeedDeps | + packages.NeedTypes | + packages.NeedTypesInfo | + packages.NeedSyntax, + Tests: true, + } + + pkgs, pkgsErr = packages.Load(cfg, ".") + }) + return pkgs, pkgsErr +} + +func (s *set[T]) getInfo(file string) *info { + if s.info != nil { + return s.info + } + info := &info{ + t: s.t, + tcType: reflect.TypeOf(new(T)).Elem(), + calls: make(map[token.Position]*callInfo), + patches: make(map[ast.Node]ast.Expr), + } + s.info = info + + t := s.t + + pkgs, pkgsErr = initPackages() + if pkgsErr != nil { + t.Fatalf("load: %v\n", pkgsErr) + } + + // Get package under test. + f, pkg := findFileAndPackage(file, pkgs) + if f == nil { + t.Fatalf("failed to load package for file %s", file) + } + info.testPkg = pkg + + // TODO: not necessary at the moment, but this is tricky so leaving this in + // so as to not to forget how to do it. + // + // for _, p := range pkg.Types.Imports() { + // if p.Path() == "cuelang.org/go/internal/tdtest" { + // info.thisPkg = p + // } + // } + // if info.thisPkg == nil { + // t.Fatalf("could not find test package") + // } + + // Find function declaration of this test. + var fn *ast.FuncDecl + for _, d := range f.Decls { + if fd, ok := d.(*ast.FuncDecl); ok && fd.Name.Name == t.Name() { + fn = fd + } + } + if fn == nil { + t.Fatalf("could not find test %q in file %q", t.Name(), file) + } + + // Find CompositLit table used for the test: + // - find call to which CompositLit was passed, + a := info.findCalls(fn.Body, "New", "Run") + if len(a) != 1 { + // TODO: allow more than one. + t.Fatalf("only one Run or New function allowed per test") + } + + // - analyse second argument of call, + call := a[0].ast + fset := info.testPkg.Fset + ti := info.testPkg.TypesInfo + ident, ok := call.Args[1].(*ast.Ident) + if !ok { + t.Fatalf("%v: arg 2 of %s must be a reference to the table", + fset.Position(call.Args[1].Pos()), a[0].funcName) + } + def := ti.Uses[ident] + pos := def.Pos() + + // - locate the CompositLit in the AST based on position. + v, ok := findVar(pos, f).(*ast.CompositeLit) + if !ok { + // generics should avoid this. + t.Fatalf("expected composite literal, found %T", v) + } + info.table = v + + // Find and index assertion calls. + a = info.findCalls(fn.Body, "Equal") + for _, x := range a { + info.initFieldRef(x, f) + } + + return info +} + +// initFieldRef updates c with information about the field referenced +// in its corresponding call: +// - name of the field +// - indexes the field based on filename and line number. +func (i *info) initFieldRef(c *callInfo, f *ast.File) { + call := c.ast + t := i.t + info := i.testPkg.TypesInfo + fset := i.testPkg.Fset + pos := fset.Position(call.Pos()) + + sel, ok := call.Args[1].(*ast.SelectorExpr) + s := info.Selections[sel] + if !ok || s == nil || s.Kind() != types.FieldVal { + t.Fatalf("%v: arg 2 of %s must be a reference to a test case field", + fset.Position(call.Args[1].Pos()), c.funcName) + } + + obj := s.Obj() + c.fieldName = obj.Name() + if _, ok := i.tcType.FieldByName(c.fieldName); !ok { + t.Fatalf("%v: could not find field %s", + fset.Position(obj.Pos()), c.fieldName) + } + + pos.Column = 0 + pos.Offset = 0 + i.calls[pos] = c +} + +// findFileAndPackage locates the ast.File and package within the given slice +// of packages, in which the given file is located. +func findFileAndPackage(path string, pkgs []*packages.Package) (*ast.File, *packages.Package) { + for _, p := range pkgs { + for i, gf := range p.GoFiles { + if gf == path { + return p.Syntax[i], p + } + } + } + return nil, nil +} + +const ( + typeT = "*cuelang.org/go/internal/tdtest.T" + tdtestParen = `("cuelang.org/go/internal/tdtest")` +) + +// findCalls finds all call expressions within a given block for functions +// or methods defined within the tdtest package. +func (i *info) findCalls(block *ast.BlockStmt, names ...string) []*callInfo { + var a []*callInfo + ast.Inspect(block, func(n ast.Node) bool { + c, ok := n.(*ast.CallExpr) + if !ok { + return true + } + sel, ok := c.Fun.(*ast.SelectorExpr) + if !ok { + return true + } + + // TODO: also test package. It would be better to test the equality + // using the information in the types.Info/packages to ensure that + // we really got the right function. + info := i.testPkg.TypesInfo + for _, name := range names { + if sel.Sel.Name == name { + if info.TypeOf(sel.X).String() == typeT { + } else if ident, ok := sel.X.(*ast.Ident); !ok { + return true // Run method. + } else if id, ok := info.Uses[ident].(*types.PkgName); ok && strings.Contains(id.String(), tdtestParen) { + } else { + return true + } + ci := &callInfo{ + funcName: name, + ast: c, + } + a = append(a, ci) + return true + } + } + + return true + }) + return a +} + +func findVar(pos token.Pos, n ast.Node) (ret ast.Expr) { + ast.Inspect(n, func(n ast.Node) bool { + if as, ok := n.(*ast.AssignStmt); ok { + for i, v := range as.Lhs { + if v.Pos() == pos { + ret = as.Rhs[i] + } + } + return false + } + return true + }) + return ret +} + +func (s *set[TC]) update() { + info := s.info + + t := s.t + fset := info.testPkg.Fset + + file := fset.Position(info.table.Pos()).Filename + var f *ast.File + for i, gof := range info.testPkg.GoFiles { + if gof == file { + f = info.testPkg.Syntax[i] + } + } + if f == nil { + t.Fatalf("file %s not in package", file) + } + + // TODO: use text-based insertion instead: + // - sort insertions and replacements on position in descending order. + // - substitute textually. + // + // We are using Apply because this is supposed to give better handling of + // comments. In practice this only works marginally better than not handling + // positions at all. Probably a lost cause. + astutil.Apply(f, func(c *astutil.Cursor) bool { + n := c.Node() + + switch x := info.patches[n]; x.(type) { + case nil: + case *ast.KeyValueExpr: + for { + c.InsertAfter(x) + x = info.patches[x] + if x == nil { + break + } + } + default: + c.Replace(x) + } + return true + }, nil) + + // TODO: use tmp files? + w, err := os.Create(file) + if err != nil { + t.Fatal(err) + } + defer w.Close() + + err = format.Node(w, fset, f) + if err != nil { + t.Fatal(err) + } +} + +func (t *T) updateField(info *info, ci *callInfo, newValue any) { + info.needsUpdate = true + + fset := info.testPkg.Fset + + e, ok := info.table.Elts[t.iter].(*ast.CompositeLit) + if !ok { + t.Fatalf("not a composite literal") + } + + isZero := false + var value ast.Expr + switch x := reflect.ValueOf(newValue); x.Kind() { + default: + s := fmt.Sprint(x) + x = reflect.ValueOf(s) + fallthrough + case reflect.String: + s := x.String() + isZero = s == "" + if !strings.ContainsRune(s, '`') && !isZero { + s = fmt.Sprintf("`%s`", s) + } else { + s = strconv.Quote(s) + } + value = &ast.BasicLit{Kind: token.STRING, Value: s} + case reflect.Bool: + if b := x.Bool(); b { + value = &ast.BasicLit{Kind: token.IDENT, Value: "true"} + } else { + value = &ast.BasicLit{Kind: token.IDENT, Value: "false"} + isZero = true + } + case reflect.Int, reflect.Int64, reflect.Int32, reflect.Int8: + i := x.Int() + value = &ast.BasicLit{Kind: token.INT, + Value: strconv.FormatInt(i, 10)} + isZero = i == 0 + case reflect.Uint, reflect.Uint64, reflect.Uint32, reflect.Uint8: + i := x.Uint() + value = &ast.BasicLit{Kind: token.INT, + Value: strconv.FormatUint(i, 10)} + isZero = i == 0 + } + + for _, x := range e.Elts { + kv, ok := x.(*ast.KeyValueExpr) + if !ok { + t.Fatalf("%v: elements must be key value pairs", + fset.Position(kv.Pos())) + } + ident, ok := kv.Key.(*ast.Ident) + if !ok { + t.Fatalf("%v: key must be an identifier", + fset.Position(kv.Pos())) + } + if ident.Name == ci.fieldName { + info.patches[kv.Value] = value + return + } + } + + if !isZero { + kv := &ast.KeyValueExpr{ + Key: &ast.Ident{Name: ci.fieldName}, + Value: value, + } + if len(e.Elts) > 0 { + var key ast.Node = e.Elts[len(e.Elts)-1] + old := info.patches[key] + if old != nil { + info.patches[kv] = old + } + info.patches[key] = kv + } else { + info.patches[e] = &ast.CompositeLit{Elts: []ast.Expr{kv}} + } + } +}