diff options
Diffstat (limited to 'internal/enumtest')
| -rw-r--r-- | internal/enumtest/enumtest.go | 54 | ||||
| -rw-r--r-- | internal/enumtest/enumtest_test.go | 31 |
2 files changed, 85 insertions, 0 deletions
diff --git a/internal/enumtest/enumtest.go b/internal/enumtest/enumtest.go new file mode 100644 index 0000000..cccb09c --- /dev/null +++ b/internal/enumtest/enumtest.go @@ -0,0 +1,54 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +// Package enumtest reads the constants of an enum type from Go source, so +// a test can check that every value is handled - and fails when a value is +// added without the code that needs it. Only tests import it. +package enumtest + +import ( + "fmt" + "go/ast" + "go/parser" + "go/token" +) + +// Names returns, in declaration order, the constants declared with type typ +// in the Go file at path, including those an iota block types implicitly. +// A type with no constants there is an error, so a renamed file or type +// cannot make a test pass by finding nothing. +func Names(path, typ string) ([]string, error) { + f, err := parser.ParseFile(token.NewFileSet(), path, nil, 0) + if err != nil { + return nil, err + } + var names []string + for _, decl := range f.Decls { + g, ok := decl.(*ast.GenDecl) + if !ok || g.Tok != token.CONST { + continue + } + inType := false + for _, spec := range g.Specs { + vs, ok := spec.(*ast.ValueSpec) + if !ok { + continue + } + switch { + case vs.Type != nil: + id, ok := vs.Type.(*ast.Ident) + inType = ok && id.Name == typ + case len(vs.Values) > 0: + inType = false + } + if inType { + for _, n := range vs.Names { + names = append(names, n.Name) + } + } + } + } + if len(names) == 0 { + return nil, fmt.Errorf("%s: no constants of type %s", path, typ) + } + return names, nil +} diff --git a/internal/enumtest/enumtest_test.go b/internal/enumtest/enumtest_test.go new file mode 100644 index 0000000..6a7bd69 --- /dev/null +++ b/internal/enumtest/enumtest_test.go @@ -0,0 +1,31 @@ +// SPDX-License-Identifier: GPL-3.0-or-later + +package enumtest + +import ( + "os" + "path/filepath" + "reflect" + "testing" +) + +// TestNames: typed constants are found in declaration order, those an iota +// block types implicitly included; untyped constants and other types are +// not. +func TestNames(t *testing.T) { + src := "package x\n\ntype K int\ntype Other int\n\nconst (\n\tA K = iota\n\tB\n\tC\n)\n\nconst (\n\tX Other = iota\n\tY\n)\n\nconst Z K = 9\n\nconst (\n\tP K = 1\n\tQ = 2\n)\n" + path := filepath.Join(t.TempDir(), "x.go") + if err := os.WriteFile(path, []byte(src), 0o644); err != nil { + t.Fatal(err) + } + got, err := Names(path, "K") + if err != nil { + t.Fatal(err) + } + if want := []string{"A", "B", "C", "Z", "P"}; !reflect.DeepEqual(got, want) { + t.Errorf("Names = %v, want %v", got, want) + } + if _, err := Names(path, "Missing"); err == nil { + t.Error("a type with no constants should be an error") + } +} |
