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
19 changes: 0 additions & 19 deletions internal/astgen/format.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,27 +7,8 @@ import (
"go/format"
"go/printer"
"go/token"

"golang.org/x/tools/imports"
)

// FormatFile formats an AST file and processes imports
func FormatFile(filePath string, f *ast.File) (string, error) {
var buf bytes.Buffer
if err := printer.Fprint(&buf, token.NewFileSet(), f); err != nil {
return "", fmt.Errorf("formatting error: %v", err)
}
out, err := imports.Process(filePath, buf.Bytes(), &imports.Options{
Fragment: true,
AllErrors: true,
Comments: true,
})
if err != nil {
return "", fmt.Errorf("formatting error: %v", err)
}
return string(bytes.ReplaceAll(out, []byte("\n}\nfunc "), []byte("\n}\n\nfunc "))), nil
}

// Format converts an AST node to formatted Go source code
func Format(node ast.Node) string {
var buf bytes.Buffer
Expand Down
6 changes: 3 additions & 3 deletions internal/fakeserver/fakeserver.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,11 @@ package fakeserver
import (
"bytes"
"fmt"
"go/format"
"text/template"

"github.com/maxbrunsfeld/counterfeiter/v6/generator"
"golang.org/x/tools/go/packages"
"golang.org/x/tools/imports"
)

// Config holds the configuration for generating a fake server.
Expand Down Expand Up @@ -92,9 +92,9 @@ func Generate(config Config, pl []*packages.Package) (*Files, error) {
}); err != nil {
return nil, fmt.Errorf("executing main template: %w", err)
}
mainBytes, err := imports.Process("main.go", mainBuf.Bytes(), nil)
mainBytes, err := format.Source(mainBuf.Bytes())
if err != nil {
return nil, fmt.Errorf("goimports main.go: %w", err)
return nil, fmt.Errorf("formatting main.go: %w", err)
}

return &Files{
Expand Down
20 changes: 20 additions & 0 deletions internal/generate/file.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,13 @@ import (
"github.com/typelate/muxt/internal/load"
)

// File is one generated Go file: the package it is written into, the
// loaded packages the types it names resolve against, and the imports its
// declarations register as they are built.
//
// Every generated file has a File of its own. The imports a file declares
// are then the ones something in it registered, so a file never carries a
// package another file needed.
type File struct {
fileSet *token.FileSet
typesCache map[string]*types.Package
Expand Down Expand Up @@ -53,6 +60,19 @@ func newFile(filePath string, fileSet *token.FileSet, list []*packages.Package)
return file, nil
}

// sibling returns a File for another generated file in the same package:
// the same loaded packages, and imports of its own.
func (file *File) sibling() *File {
return &File{
fileSet: file.fileSet,
typesCache: file.typesCache,
files: file.files,
packages: file.packages,
outPkg: file.outPkg,
packageIdentifiers: make(map[string]string),
}
}

func (file *File) Package(path string) (*packages.Package, bool) {
return load.PackageWithPath(file.packages, path)
}
Expand Down
115 changes: 115 additions & 0 deletions internal/generate/format.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,115 @@
package generate

import (
"bytes"
"cmp"
"fmt"
"go/ast"
"go/format"
"go/printer"
"go/token"
"slices"
"strconv"
"strings"
)

// formatFile prints a generated file and formats it with go/format.
//
// The imports are the file's own: each file registers the packages its
// declarations reference as it builds them, so there is nothing to add or
// remove. They are laid out the way gofmt users expect, the standard
// library first and every other path after it, each group sorted and set
// apart by a blank line.
func formatFile(filePath string, f *ast.File) (string, error) {
var imports []*ast.ImportSpec
decls := f.Decls[:0:0]
for _, decl := range f.Decls {
if gen, ok := decl.(*ast.GenDecl); ok && gen.Tok == token.IMPORT {
for _, spec := range gen.Specs {
imports = append(imports, spec.(*ast.ImportSpec))
}
continue
}
decls = append(decls, decl)
}
body := *f
body.Decls = decls

var buf bytes.Buffer
fmt.Fprintf(&buf, "package %s\n\n", f.Name.Name)
if err := writeImports(&buf, imports); err != nil {
return "", fmt.Errorf("formatting %s: %w", filePath, err)
}
var rest bytes.Buffer
if err := printer.Fprint(&rest, token.NewFileSet(), &body); err != nil {
return "", fmt.Errorf("formatting %s: %w", filePath, err)
}
// The printed file repeats the package clause written above.
_, afterClause, _ := bytes.Cut(rest.Bytes(), []byte("\n"))
buf.Write(afterClause)

out, err := format.Source(buf.Bytes())
if err != nil {
return "", fmt.Errorf("formatting %s: %w", filePath, err)
}
return string(bytes.ReplaceAll(out, []byte("\n}\nfunc "), []byte("\n}\n\nfunc "))), nil
}

// writeImports writes an import declaration for specs, grouped by
// importGroup and sorted by path within each group, with a blank line
// between groups.
func writeImports(buf *bytes.Buffer, specs []*ast.ImportSpec) error {
type entry struct {
name, path string
group int
}
entries := make([]entry, 0, len(specs))
for _, spec := range specs {
path, err := strconv.Unquote(spec.Path.Value)
if err != nil {
return err
}
e := entry{path: path, group: importGroup(path)}
if spec.Name != nil {
e.name = spec.Name.Name
}
entries = append(entries, e)
}
slices.SortFunc(entries, func(a, b entry) int {
return cmp.Or(cmp.Compare(a.group, b.group), cmp.Compare(a.path, b.path), cmp.Compare(a.name, b.name))
})
entries = slices.Compact(entries)

line := func(e entry) string {
if e.name != "" {
return e.name + " " + strconv.Quote(e.path)
}
return strconv.Quote(e.path)
}
switch len(entries) {
case 0:
return nil
case 1:
fmt.Fprintf(buf, "import %s\n\n", line(entries[0]))
return nil
}
buf.WriteString("import (\n")
for i, e := range entries {
if i > 0 && e.group != entries[i-1].group {
buf.WriteString("\n")
}
buf.WriteString("\t" + line(e) + "\n")
}
buf.WriteString(")\n\n")
return nil
}

// importGroup orders an import path: the standard library, whose paths have
// no dot in their first element, before everything else.
func importGroup(path string) int {
first, _, _ := strings.Cut(path, "/")
if strings.Contains(first, ".") {
return 1
}
return 0
}
73 changes: 73 additions & 0 deletions internal/generate/format_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
package generate

import (
"go/ast"
"go/token"
"testing"

"github.com/typelate/muxt/internal/astgen"
)

// TestFormatFileImports states how a generated file lays out its imports:
// the standard library first, then every path whose first element has a
// dot, each group sorted by path and set apart by a blank line -- the layout
// goimports gives a file.
func TestFormatFileImports(t *testing.T) {
spec := func(name, path string) *ast.ImportSpec {
s := &ast.ImportSpec{Path: astgen.String(path)}
if name != "" {
s.Name = ast.NewIdent(name)
}
return s
}
for _, tt := range []struct {
name string
specs []*ast.ImportSpec
want string
}{
{
name: "one import",
specs: []*ast.ImportSpec{spec("", "net/http")},
want: "package server\n\nimport \"net/http\"\n",
},
{
name: "the standard library before the rest",
specs: []*ast.ImportSpec{
spec("", "github.com/example/app/models"),
spec("", "net/http"),
spec("v2", "example.com/lib/v2"),
spec("", "bytes"),
spec("", "server/internal/data"),
spec("", "server/v1.2/data"),
},
want: "package server\n\nimport (\n\t\"bytes\"\n\t\"net/http\"\n\t\"server/internal/data\"\n\t\"server/v1.2/data\"\n\n\tv2 \"example.com/lib/v2\"\n\t\"github.com/example/app/models\"\n)\n",
},
{
name: "a repeated import once",
specs: []*ast.ImportSpec{spec("", "net/http"), spec("", "net/http")},
want: "package server\n\nimport \"net/http\"\n",
},
{
name: "no imports",
want: "package server\n",
},
} {
t.Run(tt.name, func(t *testing.T) {
f := &ast.File{Name: ast.NewIdent("server")}
if len(tt.specs) > 0 {
decl := &ast.GenDecl{Tok: token.IMPORT}
for _, s := range tt.specs {
decl.Specs = append(decl.Specs, s)
}
f.Decls = []ast.Decl{decl}
}
got, err := formatFile("routes.go", f)
if err != nil {
t.Fatal(err)
}
if got != tt.want {
t.Errorf("formatFile =\n%s\nwant\n%s", got, tt.want)
}
})
}
}
9 changes: 6 additions & 3 deletions internal/generate/html.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,7 +75,10 @@ func executeHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def
},
}

if handlerFunc.Body.List, err = appendParseArgumentStatements(handlerFunc.Body.List, def, file, resultType, sig, def.Arguments, nil, resultDataIdent, config, def.CallExpression(), func(s string) *ast.BlockStmt {
// Parsing rewrites the call's arguments to the locals it declares,
// so it works on a copy and the definition stays as resolved.
call := cloneCall(def.CallExpression())
if handlerFunc.Body.List, err = appendParseArgumentStatements(handlerFunc.Body.List, def, file, resultType, sig, def.Arguments, nil, resultDataIdent, config, call, func(s string) *ast.BlockStmt {
errBlock := appendTemplateDataError(file, resultDataIdent, astgen.ErrorsNew(file, astgen.String(s)))
errBlock.List = append(errBlock.List, assignTemplateDataErrStatusCode(file, resultDataIdent, http.StatusBadRequest))
return errBlock
Expand All @@ -98,7 +101,7 @@ func executeHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def
Tok: token.VAR,
Specs: []ast.Spec{&ast.ValueSpec{Names: []*ast.Ident{ast.NewIdent(guardIdent)}, Type: astgen.ExportedIdentifier(file, "", "sync/atomic", "Bool")}},
}})
callArgs := slices.Clone(def.CallExpression().Args)
callArgs := slices.Clone(call.Args)
callArgs[execIdx] = closure
if config.Logger {
handlerFunc.Body.List = append(handlerFunc.Body.List, logDebugStatement(file, "handling request", def.RawPattern()))
Expand Down Expand Up @@ -130,7 +133,7 @@ func executeHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def
Sel: ast.NewIdent(TemplateDataFieldIdentifierResult),
}, sig, def.FunctionIdentifier().Name, &ast.CallExpr{
Fun: callFun,
Args: slices.Clone(def.CallExpression().Args),
Args: slices.Clone(call.Args),
}, errBody)
if err != nil {
return nil, err
Expand Down
Loading