Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,7 @@ numeric value as the key.
| Flag | What it does |
|---|---|
| `-values` | Adds `Values() []string`, which [ent](https://entgo.io/docs/schema-fields/#enum-fields) uses for enum fields. |
| `-validate` | Adds `Validate() error`, which returns an error when the value is not a declared constant. |
| `-flag.value` | Adds `Set(string) error` so the type satisfies `flag.Value`. |
| `-pflag.value` | Adds `Set` and `Type() string` so the type satisfies [pflag.Value](https://pkg.go.dev/github.com/spf13/pflag#Value). `Type` returns all names joined by `\|`. |
| `-typederrors` | Wraps conversion errors with `enumerrs.ErrValueInvalid`. See [Typed errors](#typed-errors). |
Expand Down Expand Up @@ -167,7 +168,7 @@ const (

## Typed errors

With `-typederrors`, `TString()` and the unmarshal methods return an error that matches
With `-typederrors`, `TString()`, `Validate()` and the unmarshal methods return an error that matches
`enumerrs.ErrValueInvalid` under `errors.Is`. The message still names the bad input.

```go
Expand Down
20 changes: 13 additions & 7 deletions endtoend_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ func TestEndToEnd(t *testing.T) {
// Names are known to be ASCII and long enough.
var typeName string
var transformNameMethod string
var useTypedErrors bool
var extraFlags []string

switch name {
case "transform_snake.go":
Expand Down Expand Up @@ -121,19 +121,27 @@ func TestEndToEnd(t *testing.T) {
case "typedErrors.go":
typeName = "TypedErrorsValue"
transformNameMethod = "noop"
useTypedErrors = true
extraFlags = []string{"-typederrors", "-values"}
case "validate.go":
typeName = "Color"
transformNameMethod = "noop"
extraFlags = []string{"-validate"}
case "validateTypedErrors.go":
typeName = "Priority"
transformNameMethod = "noop"
extraFlags = []string{"-validate", "-typederrors"}
default:
typeName = fmt.Sprintf("%c%s", name[0]+'A'-'a', name[1:len(name)-len(".go")])
transformNameMethod = "noop"
}

stringerCompileAndRun(t, dir, stringer, typeName, name, transformNameMethod, useTypedErrors)
stringerCompileAndRun(t, dir, stringer, typeName, name, transformNameMethod, extraFlags...)
}
}

// stringerCompileAndRun runs stringer for the named file and compiles and
// runs the target binary in directory dir. That binary will panic if the String method is incorrect.
func stringerCompileAndRun(t *testing.T, dir, stringer, typeName, fileName, transformNameMethod string, useTypedErrors bool) {
func stringerCompileAndRun(t *testing.T, dir, stringer, typeName, fileName, transformNameMethod string, extraFlags ...string) {
t.Logf("run: %s %s\n", fileName, typeName)
source := filepath.Join(dir, fileName)
err := copy(source, filepath.Join("testdata", fileName))
Expand All @@ -143,9 +151,7 @@ func stringerCompileAndRun(t *testing.T, dir, stringer, typeName, fileName, tran
stringSource := filepath.Join(dir, typeName+"_string.go")
// Run stringer in temporary directory.
args := []string{"-type", typeName, "-output", stringSource, "-transform", transformNameMethod}
if useTypedErrors {
args = append(args, "-typederrors", "-values")
}
args = append(args, extraFlags...)
args = append(args, source)
err = run(stringer, args...)
if err != nil {
Expand Down
21 changes: 21 additions & 0 deletions enumer.go
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,19 @@ const altStringValuesMethod = `func (%[1]s) Values() []string {
}
`

// Arguments to format are:
//
// [1]: type name
// [2]: error expression returned for a value not in the enum
const validateMethod = `// Validate returns an error if the value is not listed in the enum definition.
func (i %[1]s) Validate() error {
if !i.IsA%[1]s() {
return %[2]s
}
return nil
}
`

func (g *Generator) buildAltStringValuesMethod(typeName string) {
g.Printf("\n")
g.Printf(altStringValuesMethod, typeName)
Expand Down Expand Up @@ -228,6 +241,14 @@ func (g *Generator) buildYAMLMethods(runs [][]Value, typeName string, runsThresh
g.Printf(yamlMethods, typeName)
}

func (g *Generator) buildValidateMethod(typeName string, useTypedErrors bool) {
errorCode := fmt.Sprintf(`fmt.Errorf("%%v does not belong to %s values", i)`, typeName)
if useTypedErrors {
errorCode = fmt.Sprintf(`errors.Join(enumerrs.ErrValueInvalid, fmt.Errorf("%%v does not belong to %s values", i))`, typeName)
}
g.Printf(validateMethod, typeName, errorCode)
}

// Arguments to format are: [1]: type name
const flagValueMethodSet = `
// Set allows flag and pflag libraries to set a value dynamically.
Expand Down
21 changes: 21 additions & 0 deletions golden_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,14 @@ var goldenLinecomment = []Golden{
{"dayWithLinecomment", linecommentIn},
}

var goldenValidate = []Golden{
{"validate", dayIn},
}

var goldenValidateTypedErrors = []Golden{
{"validateTypedErrors", dayIn},
}

var goldenFlagValue = []Golden{
{"flagvalue", dayIn},
}
Expand Down Expand Up @@ -408,6 +416,19 @@ func TestGolden(t *testing.T) {
lineComment: true,
})
}
for _, test := range goldenValidate {
runGoldenTest(t, test, generateOptions{
transformMethod: "noop",
includeValidateMethod: true,
})
}
for _, test := range goldenValidateTypedErrors {
runGoldenTest(t, test, generateOptions{
transformMethod: "noop",
includeValidateMethod: true,
useTypedErrors: true,
})
}
for _, test := range goldenFlagValue {
runGoldenTest(t, test, generateOptions{
transformMethod: "noop",
Expand Down
31 changes: 18 additions & 13 deletions stringer.go
Original file line number Diff line number Diff line change
Expand Up @@ -44,19 +44,20 @@ func (af *arrayFlags) Set(value string) error {
}

type generateOptions struct {
includeJSON bool
includeYAML bool
includeSQL bool
includeText bool
includeGQLGen bool
transformMethod string
trimPrefix string
addPrefix string
lineComment bool
includeValuesMethod bool
includeFlagMethods bool
includePflagMethods bool
useTypedErrors bool
includeJSON bool
includeYAML bool
includeSQL bool
includeText bool
includeGQLGen bool
transformMethod string
trimPrefix string
addPrefix string
lineComment bool
includeValuesMethod bool
includeValidateMethod bool
includeFlagMethods bool
includePflagMethods bool
useTypedErrors bool
}

var (
Expand All @@ -75,6 +76,7 @@ func init() {
flag.BoolVar(&opts.includeText, "text", false, "if true, text marshaling methods will be generated. Default: false")
flag.BoolVar(&opts.includeGQLGen, "gqlgen", false, "if true, GraphQL marshaling methods for gqlgen will be generated. Default: false")
flag.BoolVar(&opts.includeValuesMethod, "values", false, "if true, alternative string values method will be generated. Default: false")
flag.BoolVar(&opts.includeValidateMethod, "validate", false, "if true, a Validate() error method will be generated. Default: false")
flag.BoolVar(&opts.includeFlagMethods, "flag.value", false, "if true, ensure that the enumeration type implements stdlib flag.Value interface. Default: false")
flag.BoolVar(&opts.includePflagMethods, "pflag.value", false, "if true, ensure that the enumeration type implements pflag.Value interface, see: https://pkg.go.dev/github.com/spf13/pflag#Value Default: false")
flag.StringVar(&output, "output", "", "output file name; default srcdir/<type>_enumer.go")
Expand Down Expand Up @@ -498,6 +500,9 @@ func (g *Generator) generate(typeName string, opts generateOptions) {
if opts.includeValuesMethod {
g.buildAltStringValuesMethod(typeName)
}
if opts.includeValidateMethod {
g.buildValidateMethod(typeName, opts.useTypedErrors)
}

g.buildNoOpOrderChangeDetect(runs, typeName)

Expand Down
69 changes: 69 additions & 0 deletions testdata/validate.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
// Copyright 2014 The Go Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.

// Validate() must return nil for every declared constant and a non-nil
// error naming the value for anything else.

package main

import (
"fmt"
"strings"
)

type Color int

const (
Red Color = iota
Green
Blue
)

// Validate must be usable through an interface, which is the point of it.
type validator interface {
Validate() error
}

var _ validator = Color(0)

func main() {
ck(Red, "Red")
ck(Blue, "Blue")
ck(-1, "Color(-1)")

// Positive: every declared value validates.
for _, c := range ColorValues() {
if err := c.Validate(); err != nil {
panic(fmt.Sprintf("validate.go: %v.Validate() = %v, want nil", c, err))
}
}

// Negative: below, above, and far out of range all fail and name the value.
for _, bad := range []Color{-1, 3, 99} {
err := bad.Validate()
if err == nil {
panic(fmt.Sprintf("validate.go: Color(%d).Validate() = nil, want error", int(bad)))
}
want := fmt.Sprintf("Color(%d) does not belong to Color values", int(bad))
if err.Error() != want {
panic(fmt.Sprintf("validate.go: got %q, want %q", err.Error(), want))
}
}

// Through the interface, same answers.
var v validator = Color(3)
if v.Validate() == nil || !strings.Contains(v.Validate().Error(), "Color(3)") {
panic("validate.go: interface call did not report Color(3)")
}
v = Green
if v.Validate() != nil {
panic("validate.go: interface call rejected Green")
}
}

func ck(c Color, str string) {
if fmt.Sprint(c) != str {
panic("validate.go: " + str)
}
}
101 changes: 101 additions & 0 deletions testdata/validate.golden
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@

const _DayName = "MondayTuesdayWednesdayThursdayFridaySaturdaySunday"

var _DayIndex = [...]uint8{0, 6, 13, 22, 30, 36, 44, 50}

const _DayLowerName = "mondaytuesdaywednesdaythursdayfridaysaturdaysunday"

func (i Day) String() string {
if i < 0 || i >= Day(len(_DayIndex)-1) {
return fmt.Sprintf("Day(%d)", i)
}
return _DayName[_DayIndex[i]:_DayIndex[i+1]]
}

// Validate returns an error if the value is not listed in the enum definition.
func (i Day) Validate() error {
if !i.IsADay() {
return fmt.Errorf("%v does not belong to Day values", i)
}
return nil
}

// An "invalid array index" compiler error signifies that the constant values have changed.
// Re-run the stringer command to generate them again.
func _DayNoOp() {
var x [1]struct{}
_ = x[Monday-(0)]
_ = x[Tuesday-(1)]
_ = x[Wednesday-(2)]
_ = x[Thursday-(3)]
_ = x[Friday-(4)]
_ = x[Saturday-(5)]
_ = x[Sunday-(6)]
}

var _DayValues = []Day{Monday, Tuesday, Wednesday, Thursday, Friday, Saturday, Sunday}

var _DayNameToValueMap = map[string]Day{
_DayName[0:6]: Monday,
_DayName[6:13]: Tuesday,
_DayName[13:22]: Wednesday,
_DayName[22:30]: Thursday,
_DayName[30:36]: Friday,
_DayName[36:44]: Saturday,
_DayName[44:50]: Sunday,
}

var _DayLowerNameToValueMap = map[string]Day{
_DayLowerName[0:6]: Monday,
_DayLowerName[6:13]: Tuesday,
_DayLowerName[13:22]: Wednesday,
_DayLowerName[22:30]: Thursday,
_DayLowerName[30:36]: Friday,
_DayLowerName[36:44]: Saturday,
_DayLowerName[44:50]: Sunday,
}

var _DayNames = []string{
_DayName[0:6],
_DayName[6:13],
_DayName[13:22],
_DayName[22:30],
_DayName[30:36],
_DayName[36:44],
_DayName[44:50],
}

// DayString retrieves an enum value from the enum constants string name.
// Throws an error if the param is not part of the enum.
func DayString(s string) (Day, error) {
if val, ok := _DayNameToValueMap[s]; ok {
return val, nil
}

if val, ok := _DayLowerNameToValueMap[strings.ToLower(s)]; ok {
return val, nil
}
return 0, fmt.Errorf("%s does not belong to Day values", s)
}

// DayValues returns all values of the enum
func DayValues() []Day {
return _DayValues
}

// DayStrings returns a slice of all String values of the enum
func DayStrings() []string {
strs := make([]string, len(_DayNames))
copy(strs, _DayNames)
return strs
}

// IsADay returns "true" if the value is listed in the enum definition. "false" otherwise
func (i Day) IsADay() bool {
for _, v := range _DayValues {
if i == v {
return true
}
}
return false
}
Loading
Loading