From b600df78772257ad6720018363990d3dbfd49f17 Mon Sep 17 00:00:00 2001 From: Ben Kraft Date: Fri, 10 Apr 2020 15:21:23 -0700 Subject: [PATCH] big refactor to put the codegen onto methods of an object --- example/generated.go | 12 +++--- generate/generate.go | 78 ++++++++++++++++++++------------------ generate/operation.go.tmpl | 4 +- generate/types.go | 38 +++++++++++++------ generate/types_test.go | 5 ++- 5 files changed, 78 insertions(+), 59 deletions(-) diff --git a/example/generated.go b/example/generated.go index 83376e8..c7c0621 100644 --- a/example/generated.go +++ b/example/generated.go @@ -14,6 +14,12 @@ type getViewerResponse struct { } `json:"viewer"` } +type getUserResponse struct { + User *struct { + TheirName *string `json:"theirName"` + } `json:"user"` +} + func getViewer(ctx context.Context, client *graphql.Client) (*getViewerResponse, error) { var retval getViewerResponse err := client.MakeRequest(ctx, ` @@ -26,12 +32,6 @@ query getViewer { return &retval, err } -type getUserResponse struct { - User *struct { - TheirName *string `json:"theirName"` - } `json:"user"` -} - // getUser gets the given user's name from their username. func getUser(ctx context.Context, client *graphql.Client, login string) (*getUserResponse, error) { variables := map[string]interface{}{ diff --git a/generate/generate.go b/generate/generate.go index cf36d22..9c94508 100644 --- a/generate/generate.go +++ b/generate/generate.go @@ -20,11 +20,16 @@ var tmplAbsFilename = filepath.Join(filepath.Dir(thisFilename), tmplRelFilename) var tmpl = template.Must(template.ParseFiles(tmplAbsFilename)) -type templateParams struct { +// generator is the context for the codegen process (and ends up getting passed +// to the template). +type generator struct { // The name of the package into which to generate the operation-helpers. PackageName string // The list of operations for which to generate code. Operations []operation + // The types needed for these operations. + typeMap map[string]string + schema *ast.Schema } type operation struct { @@ -38,11 +43,8 @@ type operation struct { Body string // The arguments to the operation. Args []argument - // The type-name for the operation's response type. ResponseName string - // The body of the operation's response type (e.g. struct { ... }). - ResponseType string } type argument struct { @@ -51,11 +53,27 @@ type argument struct { GraphQLName string } -func fromASTArg(arg *ast.VariableDefinition, schema *ast.Schema) (argument, error) { +func newGenerator(packageName string, schema *ast.Schema) *generator { + return &generator{ + PackageName: packageName, + typeMap: map[string]string{}, + schema: schema, + } +} + +func (g *generator) Types() string { + defs := make([]string, 0, len(g.typeMap)) + for _, def := range g.typeMap { + defs = append(defs, def) + } + return strings.Join(defs, "\n\n") +} + +func (g *generator) getArgument(arg *ast.VariableDefinition) (argument, error) { graphQLName := arg.Variable firstRest := strings.SplitN(graphQLName, "", 2) goName := strings.ToLower(firstRest[0]) + firstRest[1] - goType, err := typeForInputType(arg.Type, schema) + goType, err := g.addTypeForInputType(arg.Type) if err != nil { return argument{}, err } @@ -66,13 +84,7 @@ func fromASTArg(arg *ast.VariableDefinition, schema *ast.Schema) (argument, erro }, nil } -func reverse(slice []string) { - for left, right := 0, len(slice)-1; left < right; left, right = left+1, right-1 { - slice[left], slice[right] = slice[right], slice[left] - } -} - -func getDocComment(op *ast.OperationDefinition) string { +func (g *generator) getDocComment(op *ast.OperationDefinition) string { var commentLines []string var sourceLines = strings.Split(op.Position.Src.Input, "\n") for i := op.Position.Line - 1; i > 0; i-- { @@ -90,7 +102,7 @@ func getDocComment(op *ast.OperationDefinition) string { return strings.Join(commentLines, "\n") } -func fromASTOperation(op *ast.OperationDefinition, schema *ast.Schema) (operation, error) { +func (g *generator) addOperation(op *ast.OperationDefinition) error { // TODO: we may have to actually get the precise query text, in case we // want to be hashing it or something like that. This is a bit tricky // because gqlparser's ast doesn't provide node end-position (only @@ -105,30 +117,28 @@ func fromASTOperation(op *ast.OperationDefinition, schema *ast.Schema) (operatio args := make([]argument, len(op.VariableDefinitions)) for i, arg := range op.VariableDefinitions { var err error - args[i], err = fromASTArg(arg, schema) + args[i], err = g.getArgument(arg) if err != nil { - return operation{}, err + return err } } - // TODO: configure ResponseName format - responseName := op.Name + "Response" - typ, err := typeForOperation(responseName, op, schema) + responseName, err := g.addTypeForOperation(op) if err != nil { - return operation{}, fmt.Errorf("could not compute return-type for query: %v", err) + return err } - return operation{ + g.Operations = append(g.Operations, operation{ Type: op.Operation, Name: op.Name, - Doc: getDocComment(op), + Doc: g.getDocComment(op), // The newline just makes it format a little nicer - Body: "\n" + builder.String(), - Args: args, - + Body: "\n" + builder.String(), + Args: args, ResponseName: responseName, - ResponseType: typ, - }, nil + }) + + return nil } func Generate(config *Config) ([]byte, error) { @@ -142,21 +152,15 @@ func Generate(config *Config) ([]byte, error) { return nil, err } - operations := make([]operation, len(document.Operations)) - for i, op := range document.Operations { - operations[i], err = fromASTOperation(op, schema) - if err != nil { + g := newGenerator(config.Package, schema) + for _, op := range document.Operations { + if err = g.addOperation(op); err != nil { return nil, err } } - data := templateParams{ - PackageName: config.Package, - Operations: operations, - } - var buf bytes.Buffer - err = tmpl.Execute(&buf, data) + err = tmpl.Execute(&buf, g) if err != nil { return nil, fmt.Errorf("could not render template: %v", err) } diff --git a/generate/operation.go.tmpl b/generate/operation.go.tmpl index f8978e5..28e7b34 100644 --- a/generate/operation.go.tmpl +++ b/generate/operation.go.tmpl @@ -8,9 +8,9 @@ import ( "github.com/Khan/genql/graphql" ) -{{range .Operations}} -{{.ResponseType}} +{{.Types}} +{{range .Operations}} {{.Doc}} func {{.Name}}(ctx context.Context, client *graphql.Client{{range .Args}}, {{.GoName}} {{.GoType}}{{end}}) (*{{.ResponseName}}, error) { {{- if .Args -}} diff --git a/generate/types.go b/generate/types.go index cdb7596..34caa3a 100644 --- a/generate/types.go +++ b/generate/types.go @@ -9,32 +9,46 @@ import ( type typeBuilder struct { strings.Builder - schema *ast.Schema + *generator } -func (builder *typeBuilder) baseTypeForOperation(operation ast.Operation) *ast.Definition { +func (g *generator) baseTypeForOperation(operation ast.Operation) *ast.Definition { switch operation { case ast.Query: - return builder.schema.Query + return g.schema.Query case ast.Mutation: - return builder.schema.Mutation + return g.schema.Mutation case ast.Subscription: - return builder.schema.Subscription + return g.schema.Subscription default: panic(fmt.Sprintf("unexpected operation: %v", operation)) } } -func typeForOperation(name string, operation *ast.OperationDefinition, schema *ast.Schema) (string, error) { - builder := &typeBuilder{schema: schema} +func (g *generator) addTypeForOperation(operation *ast.OperationDefinition) (name string, err error) { + // TODO: configure ResponseName format + name = operation.Name + "Response" + + if def, ok := g.typeMap[name]; ok { + // TODO: if the name is taken, maybe try to find another? + return "", fmt.Errorf("%s already defined:\n%s", name, def) + } + + builder := &typeBuilder{generator: g} fmt.Fprintf(builder, "type %s ", name) - err := builder.writeTypedef( - builder.baseTypeForOperation(operation.Operation), operation.SelectionSet) - return builder.String(), err + err = builder.writeTypedef( + g.baseTypeForOperation(operation.Operation), operation.SelectionSet) + if err != nil { + return "", err + } + + def := builder.String() + g.typeMap[name] = def + return name, nil } -func typeForInputType(typ *ast.Type, schema *ast.Schema) (string, error) { - builder := &typeBuilder{schema: schema} +func (g *generator) addTypeForInputType(typ *ast.Type) (string, error) { + builder := &typeBuilder{generator: g} // TODO: handle non-scalar types (by passing ...something... as the // SelectionSet?) diff --git a/generate/types_test.go b/generate/types_test.go index df54d77..70f9cec 100644 --- a/generate/types_test.go +++ b/generate/types_test.go @@ -118,13 +118,14 @@ func TestTypeForOperation(t *testing.T) { t.Fatalf("got %v operations, want 1", len(queryDoc.Operations)) } - goType, err := typeForOperation("Response", queryDoc.Operations[0], schema) + g := newGenerator("test_package", schema) + name, err := g.addTypeForOperation(queryDoc.Operations[0]) if err != nil { t.Error(err) } // gofmt before comparing. - goType, err = gofmt(goType) + goType, err := gofmt(g.typeMap[name]) if err != nil { t.Error(err) }