add support for enums
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"go/format"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sort"
|
||||
"strings"
|
||||
"text/template"
|
||||
|
||||
@@ -66,6 +67,10 @@ func (g *generator) Types() string {
|
||||
for _, def := range g.typeMap {
|
||||
defs = append(defs, def)
|
||||
}
|
||||
// Make sure we have a stable order. (It's somewhat
|
||||
// arbitrary but in practice mostly alphabetical.)
|
||||
// TODO: ideally we'd do a nice semantic ordering.
|
||||
sort.Strings(defs)
|
||||
return strings.Join(defs, "\n\n")
|
||||
}
|
||||
|
||||
|
||||
+51
-21
@@ -8,6 +8,7 @@ import (
|
||||
)
|
||||
|
||||
type typeBuilder struct {
|
||||
typeName string
|
||||
strings.Builder
|
||||
*generator
|
||||
}
|
||||
@@ -31,34 +32,43 @@ func (g *generator) getTypeForOperation(operation *ast.OperationDefinition) (nam
|
||||
|
||||
if def, ok := g.typeMap[name]; ok {
|
||||
// TODO: check for and handle conflicts a better way
|
||||
return name, fmt.Errorf("%s already defined:\n%s", name, def)
|
||||
return "", fmt.Errorf("%s already defined:\n%s", name, def)
|
||||
}
|
||||
|
||||
selectionSet, err := selections(operation.SelectionSet)
|
||||
if err != nil {
|
||||
return name, err
|
||||
return "", err
|
||||
}
|
||||
|
||||
err = g.addTypeForDefinition(
|
||||
return g.addTypeForDefinition(
|
||||
name, g.baseTypeForOperation(operation.Operation), selectionSet)
|
||||
|
||||
return name, err
|
||||
}
|
||||
|
||||
func (g *generator) addTypeForDefinition(name string, typ *ast.Definition, selectionSet []selection) error {
|
||||
builder := &typeBuilder{generator: g}
|
||||
func (g *generator) addTypeForDefinition(nameOverride string, typ *ast.Definition, selectionSet []selection) (name string, err error) {
|
||||
if nameOverride != "" {
|
||||
name = nameOverride
|
||||
} else {
|
||||
// TODO: casing should be configurable
|
||||
name = lowerFirst(typ.Name)
|
||||
}
|
||||
|
||||
if _, ok := g.typeMap[name]; ok {
|
||||
return name, nil
|
||||
}
|
||||
|
||||
builder := &typeBuilder{typeName: name, generator: g}
|
||||
fmt.Fprintf(builder, "type %s ", name)
|
||||
err := builder.writeTypedef(typ, selectionSet)
|
||||
err = builder.writeTypedef(typ, selectionSet)
|
||||
if err != nil {
|
||||
return err
|
||||
return "", err
|
||||
}
|
||||
|
||||
g.typeMap[name] = builder.String()
|
||||
return nil
|
||||
return name, nil
|
||||
}
|
||||
|
||||
func (g *generator) getTypeForInputType(typ *ast.Type) (string, error) {
|
||||
builder := &typeBuilder{generator: g}
|
||||
builder := &typeBuilder{typeName: lowerFirst(typ.Name()), generator: g}
|
||||
err := builder.writeType(typ, selectionsForType(g, typ), false)
|
||||
return builder.String(), err
|
||||
}
|
||||
@@ -186,20 +196,28 @@ func (builder *typeBuilder) writeType(typ *ast.Type, selectionSet []selection, i
|
||||
builder.WriteString("*")
|
||||
}
|
||||
|
||||
_, ok := builtinTypes[typ.Name()]
|
||||
def := builder.schema.Types[typ.Name()]
|
||||
if ok || inline {
|
||||
// TODO: set inline = false for nested types
|
||||
switch def.Kind {
|
||||
case ast.Enum, ast.Union, ast.Interface:
|
||||
inline = false
|
||||
case ast.Scalar:
|
||||
// TODO: this makes no sense! refactor builtin type handling.
|
||||
inline = true
|
||||
}
|
||||
|
||||
if inline {
|
||||
return builder.writeTypedef(def, selectionSet)
|
||||
}
|
||||
|
||||
// TODO: casing should be configurable?
|
||||
name := lowerFirst(typ.Name())
|
||||
builder.WriteString(name)
|
||||
if _, ok := builder.typeMap[name]; ok {
|
||||
return nil
|
||||
// Writes a typedef elsewhere (if not already defined)
|
||||
name, err := builder.addTypeForDefinition("", def, selectionSet)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Writes a typedef elsewhere
|
||||
return builder.addTypeForDefinition(name, def, selectionSet)
|
||||
|
||||
builder.WriteString(name)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (builder *typeBuilder) writeTypedef(typedef *ast.Definition, selectionSet []selection) error {
|
||||
@@ -214,7 +232,7 @@ func (builder *typeBuilder) writeTypedef(typedef *ast.Definition, selectionSet [
|
||||
}
|
||||
builder.WriteString("}")
|
||||
return nil
|
||||
case ast.Scalar, ast.Enum:
|
||||
case ast.Scalar:
|
||||
goName := builtinTypes[typedef.Name]
|
||||
// TODO(benkraft): Handle custom scalars and enums.
|
||||
if goName == "" {
|
||||
@@ -222,6 +240,18 @@ func (builder *typeBuilder) writeTypedef(typedef *ast.Definition, selectionSet [
|
||||
}
|
||||
builder.WriteString(goName)
|
||||
return nil
|
||||
case ast.Enum:
|
||||
// All GraphQL enums have underlying type string (in the Go sense).
|
||||
builder.WriteString("string\n")
|
||||
builder.WriteString("const (\n")
|
||||
for _, val := range typedef.EnumValues {
|
||||
// TODO: casing should be configurable
|
||||
fmt.Fprintf(builder, "%s %s = \"%s\"\n",
|
||||
goConstName(val.Name+"_"+builder.typeName),
|
||||
builder.typeName, val.Name)
|
||||
}
|
||||
builder.WriteString(")\n")
|
||||
return nil
|
||||
case ast.Union, ast.Interface:
|
||||
return fmt.Errorf("not implemented: %v", typedef.Kind)
|
||||
default:
|
||||
|
||||
+57
-7
@@ -3,6 +3,7 @@ package generate
|
||||
import (
|
||||
"fmt"
|
||||
"go/format"
|
||||
"sort"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -11,6 +12,7 @@ import (
|
||||
)
|
||||
|
||||
func gofmt(src string) (string, error) {
|
||||
src = strings.TrimSpace(src)
|
||||
formatted, err := format.Source([]byte(src))
|
||||
if err != nil {
|
||||
return src, err
|
||||
@@ -19,10 +21,16 @@ func gofmt(src string) (string, error) {
|
||||
}
|
||||
|
||||
var schemaText = `
|
||||
enum Role {
|
||||
STUDENT
|
||||
TEACHER
|
||||
}
|
||||
|
||||
input UserQueryInput {
|
||||
email: String
|
||||
name: String
|
||||
id: ID
|
||||
role: Role
|
||||
}
|
||||
|
||||
type AuthMethod {
|
||||
@@ -32,6 +40,7 @@ var schemaText = `
|
||||
|
||||
type User {
|
||||
id: ID!
|
||||
roles: [Role!]
|
||||
name: String
|
||||
emails: [String!]!
|
||||
emailsOrNull: [String!]
|
||||
@@ -68,6 +77,20 @@ func TestTypeForOperation(t *testing.T) {
|
||||
}`,
|
||||
// Here on out, we use aliases, just because aliases are a lot less
|
||||
// annoying to write in Go strings than Go struct tags.
|
||||
}, {
|
||||
"QueryWithDoubleAlias",
|
||||
`{
|
||||
User: user {
|
||||
ID: id
|
||||
AlsoID: id
|
||||
}
|
||||
}`,
|
||||
`type Response struct{
|
||||
User *struct {
|
||||
ID string
|
||||
AlsoID string
|
||||
}
|
||||
}`,
|
||||
}, {
|
||||
"QueryWithSlices",
|
||||
`{
|
||||
@@ -104,6 +127,24 @@ func TestTypeForOperation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}`,
|
||||
}, {
|
||||
"QueryWithEnums",
|
||||
`{
|
||||
User: user {
|
||||
Roles: roles
|
||||
}
|
||||
}`,
|
||||
`type Response struct{
|
||||
User *struct {
|
||||
Roles []role
|
||||
}
|
||||
}
|
||||
|
||||
type role string
|
||||
const (
|
||||
studentRole role = "STUDENT"
|
||||
teacherRole role = "TEACHER"
|
||||
)`,
|
||||
}}
|
||||
|
||||
for _, test := range tests {
|
||||
@@ -130,13 +171,13 @@ func TestTypeForOperation(t *testing.T) {
|
||||
}
|
||||
|
||||
g := newGenerator("test_package", schema)
|
||||
name, err := g.getTypeForOperation(queryDoc.Operations[0])
|
||||
_, err = g.getTypeForOperation(queryDoc.Operations[0])
|
||||
if err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
|
||||
// gofmt before comparing.
|
||||
goType, err := gofmt(g.typeMap[name])
|
||||
goType, err := gofmt(g.Types())
|
||||
if err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
@@ -168,16 +209,25 @@ func TestTypeForInputType(t *testing.T) {
|
||||
`DefinedType`,
|
||||
`UserQueryInput`,
|
||||
`*userQueryInput`,
|
||||
[]string{`type userQueryInput struct {
|
||||
Email *string ` + "`json:\"email\"`" + `
|
||||
Name *string ` + "`json:\"name\"`" + `
|
||||
Id *string ` + "`json:\"id\"`" + `
|
||||
}`},
|
||||
[]string{
|
||||
`type role string
|
||||
const (
|
||||
studentRole role = "STUDENT"
|
||||
teacherRole role = "TEACHER"
|
||||
)`,
|
||||
`type userQueryInput struct {
|
||||
Email *string ` + "`json:\"email\"`" + `
|
||||
Name *string ` + "`json:\"name\"`" + `
|
||||
Id *string ` + "`json:\"id\"`" + `
|
||||
Role *role ` + "`json:\"role\"`" + `
|
||||
}`,
|
||||
},
|
||||
}}
|
||||
|
||||
for _, test := range tests {
|
||||
test := test
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
sort.Strings(test.otherTypes) // To match generator.Types()
|
||||
expectedGoCode := fmt.Sprintf(
|
||||
"type Input %s\n\n%s", test.expectedGoType,
|
||||
strings.Join(test.otherTypes, "\n\n"))
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package generate
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
@@ -26,3 +27,21 @@ func lowerFirst(s string) string {
|
||||
func upperFirst(s string) string {
|
||||
return changeFirst(s, unicode.ToUpper)
|
||||
}
|
||||
|
||||
func goConstName(s string) string {
|
||||
var prev rune
|
||||
return strings.Map(func(r rune) rune {
|
||||
var ret rune
|
||||
if prev == 0 && r == '_' {
|
||||
return '_' // still treat next char as first
|
||||
} else if r == '_' {
|
||||
ret = -1
|
||||
} else if prev == '_' {
|
||||
ret = unicode.ToUpper(r)
|
||||
} else {
|
||||
ret = unicode.ToLower(r)
|
||||
}
|
||||
prev = r
|
||||
return ret
|
||||
}, s)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
package generate
|
||||
|
||||
import "testing"
|
||||
|
||||
type test struct {
|
||||
name string
|
||||
in string
|
||||
out string
|
||||
}
|
||||
|
||||
func testStringFunc(t *testing.T, f func(string) string, tests []test) {
|
||||
for _, test := range tests {
|
||||
test := test
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
got := f(test.in)
|
||||
if got != test.out {
|
||||
t.Errorf("got %#v want %#v", got, test.out)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLowerFirst(t *testing.T) {
|
||||
tests := []test{
|
||||
{"Empty", "", ""},
|
||||
{"SingleLower", "l", "l"},
|
||||
{"SingleUpper", "L", "l"},
|
||||
{"SingleUnicodeLower", "ļ", "ļ"},
|
||||
{"SingleUnicodeUpper", "Ļ", "ļ"},
|
||||
{"LongerLower", "lasdf", "lasdf"},
|
||||
{"LongerUpper", "Lasdf", "lasdf"},
|
||||
{"LongerUnicodeLower", "ļasdf", "ļasdf"},
|
||||
{"LongerUnicodeUpper", "Ļasdf", "ļasdf"},
|
||||
}
|
||||
|
||||
testStringFunc(t, lowerFirst, tests)
|
||||
}
|
||||
|
||||
func TestUpperFirst(t *testing.T) {
|
||||
tests := []test{
|
||||
{"Empty", "", ""},
|
||||
{"SingleLower", "l", "L"},
|
||||
{"SingleUpper", "L", "L"},
|
||||
{"SingleUnicodeLower", "ļ", "Ļ"},
|
||||
{"SingleUnicodeUpper", "Ļ", "Ļ"},
|
||||
{"LongerLower", "lasdf", "Lasdf"},
|
||||
{"LongerUpper", "Lasdf", "Lasdf"},
|
||||
{"LongerUnicodeLower", "ļasdf", "Ļasdf"},
|
||||
{"LongerUnicodeUpper", "Ļasdf", "Ļasdf"},
|
||||
}
|
||||
|
||||
testStringFunc(t, upperFirst, tests)
|
||||
}
|
||||
|
||||
func TestGoConstName(t *testing.T) {
|
||||
tests := []test{
|
||||
{"Empty", "", ""},
|
||||
{"AllCaps", "ASDF", "asdf"},
|
||||
{"AllCapsWithUnderscore", "ASDF_GH", "asdfGh"},
|
||||
{"JustUnderscore", "_", "_"},
|
||||
{"LeadingUnderscore", "_ASDF", "_asdf"},
|
||||
}
|
||||
|
||||
testStringFunc(t, goConstName, tests)
|
||||
}
|
||||
Reference in New Issue
Block a user