package main import ( "fmt" "go/format" "os" "path/filepath" "strings" "text/template" ) func cmdEnumgen(args []string) { dirs := args if len(dirs) == 0 { dirs = []string{"."} } for _, dir := range dirs { if err := enumgenDir(dir); err != nil { fmt.Fprintf(os.Stderr, "atlas9 enumgen: %s: %v\n", dir, err) os.Exit(1) } } } func enumgenDir(dir string) error { entries, err := os.ReadDir(dir) if err != nil { return err } for _, entry := range entries { if entry.IsDir() { continue } name := entry.Name() if !strings.HasSuffix(name, ".go") { continue } if strings.HasSuffix(name, "_atlas_enum.go") { continue } if strings.HasSuffix(name, "_test.go") { continue } path := filepath.Join(dir, name) if err := enumgenFile(dir, path); err != nil { return fmt.Errorf("%s: %w", name, err) } } return nil } func enumgenFile(dir, path string) error { data, err := os.ReadFile(path) if err != nil { return err } lines := strings.Split(string(data), "\n") var pkgName string var directives []enumDirective for _, line := range lines { line = strings.TrimSpace(line) if pkgName == "" && strings.HasPrefix(line, "package ") { fields := strings.Fields(line) if len(fields) >= 2 { pkgName = fields[1] } } if rest, ok := strings.CutPrefix(line, "//atlas:enum "); ok { parts := strings.Fields(rest) if len(parts) < 2 { continue } directives = append(directives, enumDirective{ TypeName: parts[0], Values: parts[1:], }) } } for _, d := range directives { d.Package = pkgName outName := strings.ToLower(d.TypeName) + "_atlas_enum.go" outPath := filepath.Join(dir, outName) if err := generateEnum(outPath, d); err != nil { return err } fmt.Printf("atlas9 enumgen: wrote %s\n", outPath) } return nil } type enumDirective struct { Package string TypeName string Values []string } type enumData struct { Package string Type string TypeLower string Values []string } var enumTmpl = template.Must(template.New("enum").Parse( `// Code generated by enumgen. DO NOT EDIT. package {{.Package}} import ( "database/sql/driver" "encoding/json" "fmt" ) type {{.Type}} struct{ val string } var ( {{.Type}}Unknown = {{.Type}}{} {{- range .Values}} {{$.Type}}{{.}} = {{$.Type}}{"{{.}}"} {{- end}} ) var all{{.Type}} = []{{.Type}}{ {{.Type}}Unknown, {{- range .Values}} {{$.Type}}{{.}}, {{- end}} } var {{.TypeLower}}ByVal = map[string]{{.Type}}{ "": {{.Type}}Unknown, {{- range .Values}} "{{.}}": {{$.Type}}{{.}}, {{- end}} } func Parse{{.Type}}(s string) ({{.Type}}, error) { v, ok := {{.TypeLower}}ByVal[s] if !ok { return {{.Type}}{}, fmt.Errorf("invalid {{.Type}} %q", s) } return v, nil } func (s {{.Type}}) Valid() bool { _, ok := {{.TypeLower}}ByVal[s.val]; return ok } func (s {{.Type}}) Values() []{{.Type}} { return all{{.Type}} } func (s {{.Type}}) String() string { return s.val } func (s {{.Type}}) FilterValues() []string { strs := make([]string, len(all{{.Type}})-1) for i, v := range all{{.Type}}[1:] { strs[i] = v.String() } return strs } func (s {{.Type}}) MarshalJSON() ([]byte, error) { return json.Marshal(s.val) } func (s *{{.Type}}) UnmarshalJSON(b []byte) error { var str string if err := json.Unmarshal(b, &str); err != nil { return err } v, err := Parse{{.Type}}(str) if err != nil { return err } *s = v return nil } func (s {{.Type}}) Value() (driver.Value, error) { return s.val, nil } func (s *{{.Type}}) Scan(src any) error { str, ok := src.(string) if !ok { return fmt.Errorf("{{.Type}}.Scan: expected string, got %T", src) } v, err := Parse{{.Type}}(str) if err != nil { return err } *s = v return nil } `)) func generateEnum(outPath string, d enumDirective) error { data := enumData{ Package: d.Package, Type: d.TypeName, TypeLower: strings.ToLower(d.TypeName[:1]) + d.TypeName[1:], Values: d.Values, } var buf strings.Builder if err := enumTmpl.Execute(&buf, data); err != nil { return err } formatted, err := format.Source([]byte(buf.String())) if err != nil { return fmt.Errorf("format %s: %w\n---\n%s", outPath, err, buf.String()) } return os.WriteFile(outPath, formatted, 0644) }