diff --git a/CLAUDE.md b/CLAUDE.md index ea591ba3..b752f86b 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -19,17 +19,25 @@ ## Architecture Overview ``` -Template Name (with route pattern and method call) +go list (golang.org/x/tools/go/packages) ./internal/load + ↓ load.Package, load.GenerateSource, load.RoutesSource +source.Package: types + templates variables ./internal/source + ↓ muxt.Definitions, muxt.ResolveCall +Resolved routes (muxt.Definition) ./internal/muxt ↓ -Parser (./internal/muxt/parse) - ↓ -Type Checker (go/types ./internal/analysis) - ↓ -Generator (./internal/muxt/generate) - ↓ -HTTP Handler Code +Generated files / check reports / mutations ./internal/generate, ./internal/analysis, ./internal/mutation ``` +**The package load stops at `internal/load`.** It is the only package that +calls `packages.Load` (the go command, seconds per run). It hydrates a +command's configuration into a `source.Package`: plain data holding the +package's types, and each templates variable's template set, functions, +definitions and ExecuteTemplate calls. Everything below it reads only that, +so it can be tested with inputs built in memory: `internal/typestest` type +checks Go source against stub standard library packages in microseconds, and +`internal/load/loadtest` builds a loaded package from it for tests that go +through `internal/load`. Use them rather than loading a module in unit tests. + **Key concept:** Muxt reads template names like `"GET /{id} GetUser(ctx, id)"` and generates `http.Handler` implementations that: - Parse URL parameters to the correct Go types - Call the receiver method with parsed args @@ -82,10 +90,26 @@ go test ./... ### 4. Implement Changes -Update the generator code in order: -1. `internal/muxt/` — Core generation logic -2. `internal/source/` — AST helpers (if needed) -3. `internal/cli/` — CLI handling (if needed) +Update the code in order: +1. `internal/muxt/` — Route name parsing and call resolution +2. `internal/generate/` or `internal/analysis/` — What is written or reported +3. `internal/load/` — Only if a run needs something new from the loaded packages +4. `internal/cli/` — CLI handling (if needed) + +Before adding an integration script, see whether a unit test can state it, +at the layer that owns the behavior: +- **Flags:** `internal/cli/configurations_test.go` states what a command line + parses into, and which command lines are rejected, without loading anything. +- **What a command does with a valid configuration:** + `internal/{generate,analysis,mutation}/testdata/*.txtar` snapshot generated + files, check reports and mutation dry runs from in-memory packages in + milliseconds. Each archive's configuration is a literal in that package's + `snapshots_test.go`, copied from the command line case it stands for. + Rewrite snapshots with `go test ./internal/generate -run TestSnapshots -update` + and review the diff. + +Integration scripts are for what needs the go command: generated code +compiling and serving requests, and files on disk. ### 5. Verify Your Changes @@ -150,8 +174,13 @@ ls cmd/muxt/testdata/err_*.txt ## Key Files and Directories ### Source Code -- `internal/muxt/` — Generator logic (parse, type check, generate) -- `internal/source/` — AST analysis helpers +- `internal/load/` — Package loading (the only `go/packages` caller), hydration into `source.Package`, and load diagnostics +- `internal/source/` — The loaded package as plain data: what everything after the load reads +- `internal/muxt/` — Template name parsing and route resolution against go/types +- `internal/generate/` — Routes file generation +- `internal/analysis/` — `muxt check` and the template listings +- `internal/typestest/` — In-memory type checking against stub standard library packages, for tests +- `internal/load/loadtest/` — A loaded package built in memory, for tests that go through `internal/load` - `internal/cli/` — Command-line interface - `cmd/muxt/` — Command entry point diff --git a/internal/analysis/check.go b/internal/analysis/check.go index fc4d47b7..7a23d7ce 100644 --- a/internal/analysis/check.go +++ b/internal/analysis/check.go @@ -3,7 +3,6 @@ package analysis import ( "errors" "fmt" - "go/ast" "go/token" "go/types" "html/template" @@ -14,11 +13,10 @@ import ( "text/template/parse" "github.com/typelate/check" - "golang.org/x/tools/go/packages" - "github.com/typelate/muxt/internal/asteval" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) // executeTemplateFunc names the method the endpoint scan reports call @@ -33,26 +31,17 @@ type CheckConfiguration struct { // Check validates the package's templates and returns how many // ExecuteTemplate call sites it checked, so the caller can report the // count on success. -func Check(config CheckConfiguration, wd string, log *log.Logger, fileSet *token.FileSet, pl []*packages.Package) (int, error) { - routesPkg, ok := asteval.PackageAtFilepath(pl, wd) - if !ok { - return 0, asteval.NoPackageError(wd, pl) - } - +func Check(config CheckConfiguration, log *log.Logger, pkg source.Package) (int, error) { var errs []error totalChecked := 0 - for _, tv := range config.TemplatesVariables { - lt, err := asteval.LoadTemplates(wd, tv, pl) - if err != nil { - return totalChecked, err - } - global, ts := lt.Global, lt.HTML + for _, lt := range pkg.Variables { + global, ts := newGlobal(pkg, lt), lt.Set // Route template names are validated here so a malformed name // surfaces with its position instead of leaving the template to // be reported as merely unused below. - if _, err := muxt.Definitions(ts, tv, lt.Templates); err != nil { + if _, err := muxt.Definitions(lt); err != nil { if multiLine, ok := errors.AsType[muxt.MultiLineError](err); ok { log.Println(multiLine.MultiLineError()) log.Println() @@ -65,15 +54,15 @@ func Check(config CheckConfiguration, wd string, log *log.Logger, fileSet *token executedTemplates := make(map[string][]TemplateExecution) checkedTemplates := 0 - for c := range lt.Templates.ExecuteTemplateCalls() { + for _, c := range lt.Calls { checkedTemplates++ - templateName, dataType := c.TemplateName, c.DataType + templateName, dataType := c.Template, c.Data if config.Verbose { log.Println("checking endpoint", templateName) } - qualifier := astgen.NewTypeFormatter(routesPkg.PkgPath).Qualifier - if err := findTemplateExecution(executedTemplates, global, fileSet, qualifier, ts, c.Call, templateName, dataType); err != nil { - log.Println(fileSet.Position(c.Call.Pos()), executeTemplateFunc, strconv.Quote(templateName), types.TypeString(dataType, qualifier)) + qualifier := astgen.NewTypeFormatter(pkg.Types.Path()).Qualifier + if err := findTemplateExecution(executedTemplates, global, qualifier, ts, c.Position, templateName, dataType); err != nil { + log.Println(c.Position, executeTemplateFunc, strconv.Quote(templateName), types.TypeString(dataType, qualifier)) if checkErr, ok := errors.AsType[*check.Error](err); ok { var sb strings.Builder if detailErr := checkErr.DetailedError(&sb, qualifier); detailErr != nil { @@ -266,8 +255,8 @@ func newTemplateExecution(pos token.Position, n any, templateName string, dataTy } } -func findTemplateExecution(executedTemplates map[string][]TemplateExecution, global *check.Global, fileSet *token.FileSet, qualifier types.Qualifier, ts *template.Template, node ast.Node, templateName string, dataType types.Type) error { - executedTemplates[templateName] = append(executedTemplates[templateName], newTemplateExecution(fileSet.Position(node.Pos()), node, templateName, dataType)) +func findTemplateExecution(executedTemplates map[string][]TemplateExecution, global *check.Global, qualifier types.Qualifier, ts *template.Template, position token.Position, templateName string, dataType types.Type) error { + executedTemplates[templateName] = append(executedTemplates[templateName], newTemplateExecution(position, nil, templateName, dataType)) ts2 := ts.Lookup(templateName) if ts2 == nil { return fmt.Errorf("template %q not found", templateName) @@ -282,3 +271,15 @@ func findTemplateExecution(executedTemplates map[string][]TemplateExecution, glo } return nil } + +// newGlobal wires a check.Global for type checking a templates variable's +// templates in pkg. +func newGlobal(pkg source.Package, variable source.Variable) *check.Global { + return check.NewGlobal(pkg.Types, pkg.Fset, check.FindTreeFunc(func(name string) (*parse.Tree, bool) { + t := variable.Set.Lookup(name) + if t == nil || t.Tree == nil { + return nil, false + } + return t.Tree, true + }), check.Functions(variable.Functions)) +} diff --git a/internal/analysis/routes.go b/internal/analysis/routes.go index e98f979b..2d8b0e37 100644 --- a/internal/analysis/routes.go +++ b/internal/analysis/routes.go @@ -2,18 +2,14 @@ package analysis import ( "bytes" - "cmp" - "go/token" "go/types" "io" "maps" "slices" "strings" - "golang.org/x/tools/go/packages" - - "github.com/typelate/muxt/internal/asteval" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) type DefinitionsConfiguration struct { @@ -57,34 +53,15 @@ func (result *Routes) WriteTo(w io.Writer) (int64, error) { return io.Copy(w, &buf) } -func NewRoutes(config DefinitionsConfiguration, wd string, _ *token.FileSet, pl []*packages.Package) ([]*Routes, error) { - pkg, ok := asteval.PackageAtFilepath(pl, wd) - if !ok { - return nil, asteval.NoPackageError(wd, pl) - } - - config.PackagePath = pkg.PkgPath - config.PackageName = pkg.Name - - var receiver *types.Named - if config.ReceiverType != "" { - var err error - receiver, err = asteval.FindType(pl, cmp.Or(config.ReceiverPackage, config.PackagePath), config.ReceiverType) - if err != nil { - return nil, err - } - } - +// NewRoutes lists each templates variable's route definitions and +// functions, and the receiver's methods when a receiver type was named. +func NewRoutes(pkg source.Package, receiver *types.Named) ([]*Routes, error) { var results []*Routes - for _, tv := range config.TemplatesVariables { - lt, ts, err := asteval.HTMLTemplates(tv, pkg) - if err != nil { - return nil, err - } - functions := lt.CollectedFunctions() + for _, tv := range pkg.Variables { + functions := tv.Funcs - definitions, err := muxt.Definitions(ts, tv, lt) + definitions, err := muxt.Definitions(tv) if err != nil { return nil, err } diff --git a/internal/analysis/snapshot_test.go b/internal/analysis/snapshot_test.go new file mode 100644 index 00000000..2728a84f --- /dev/null +++ b/internal/analysis/snapshot_test.go @@ -0,0 +1,205 @@ +package analysis_test + +import ( + "bytes" + "errors" + "flag" + "fmt" + "io" + "log" + "os" + "path/filepath" + "slices" + "strings" + "testing" + + "golang.org/x/tools/txtar" + + "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/load" + "github.com/typelate/muxt/internal/load/loadtest" + "github.com/typelate/muxt/internal/muxt" +) + +var update = flag.Bool("update", false, "rewrite the want/ files of the snapshot archives in testdata") + +// templatesGo declares the templates variable for an archive that does not +// declare its own: every template file, parsed as ParseFS would. +const templatesGo = `package server + +import ( + "embed" + "html/template" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) +` + +// TestSnapshots runs an analysis for each archive in testdata and +// compares what it reports with the archive's want/ files. +// +// An archive holds a case's inputs and what the analysis reports, and +// snapshots, in snapshots_test.go, the configuration it runs with -- which +// also says which analysis runs: +// +// - Go files and .gohtml files are loaded as example.com/server by +// internal/load/loadtest -- type checked against the stub standard +// library, without the go command -- and hydrated by internal/load, as +// the command does. An archive with no templates.go gets one declaring +// the templates variable over every .gohtml file. +// - want/ files are what was reported: want/stdout.txt for a command's +// output, want/log.txt for what check logged, and want/error.txt for +// the error returned. +// +// Paths in the output are relative to the directory the package was +// written to. Run with -update to rewrite the want/ files, then read the +// diff. +func TestSnapshots(t *testing.T) { + archives, err := filepath.Glob(filepath.Join("testdata", "*.txtar")) + if err != nil { + t.Fatal(err) + } + for _, archivePath := range archives { + name := strings.TrimSuffix(filepath.Base(archivePath), ".txtar") + if !slices.ContainsFunc(snapshots, func(c snapshotCase) bool { return c.archive == name }) { + t.Errorf("testdata/%s.txtar has no configuration in snapshots", name) + } + } + for _, tt := range snapshots { + t.Run(tt.archive, func(t *testing.T) { + archivePath := filepath.Join("testdata", tt.archive+".txtar") + archive, err := txtar.ParseFile(archivePath) + if err != nil { + t.Fatal(err) + } + got := snapshot(t, tt.config, archive) + if *update { + files := slices.DeleteFunc(slices.Clone(archive.Files), func(file txtar.File) bool { + return strings.HasPrefix(file.Name, "want/") + }) + for _, name := range sortedKeys(got) { + files = append(files, txtar.File{Name: "want/" + name, Data: []byte(got[name])}) + } + archive.Files = files + if err := os.WriteFile(archivePath, txtar.Format(archive), 0o644); err != nil { + t.Fatal(err) + } + return + } + want := make(map[string]string) + for _, file := range archive.Files { + if name, ok := strings.CutPrefix(file.Name, "want/"); ok { + want[name] = string(file.Data) + } + } + for _, name := range sortedKeys(got, want) { + if got[name] != want[name] { + t.Errorf("want/%s differs (run go test -run TestSnapshots -update to rewrite):\n--- got\n%s\n--- want\n%s", name, got[name], want[name]) + } + } + }) + } +} + +func snapshot(t *testing.T, config any, archive *txtar.Archive) map[string]string { + t.Helper() + files := make(map[string]string) + for _, file := range archive.Files { + if strings.HasPrefix(file.Name, "want/") { + continue + } + files[file.Name] = string(file.Data) + } + if _, declared := files["templates.go"]; !declared { + files["templates.go"] = templatesGo + } + dir := t.TempDir() + pl := loadtest.Package(t, dir, "example.com/server", files) + relative := func(text string) string { return strings.ReplaceAll(text, dir+string(filepath.Separator), "") } + + got := make(map[string]string) + var stdout bytes.Buffer + var runErr error + switch config := config.(type) { + case analysis.CheckConfiguration: + pkg, err := load.Package(dir, pl, config.TemplatesVariables) + if err != nil { + runErr = err + break + } + var logs strings.Builder + var n int + n, runErr = analysis.Check(config, log.New(&logs, "", 0), pkg) + fmt.Fprintf(&stdout, "checked %d\n", n) + if logs.Len() > 0 { + got["log.txt"] = relative(logs.String()) + } + case analysis.DefinitionsConfiguration: + pkg, receiver, err := load.RoutesSource(dir, pl, config) + if err != nil { + runErr = err + break + } + var results []*analysis.Routes + results, runErr = analysis.NewRoutes(pkg, receiver) + for _, result := range results { + writeTo(t, &stdout, result) + } + case analysis.TemplateCallersConfiguration: + pkg, err := load.Package(dir, pl, config.TemplatesVariables) + if err != nil { + runErr = err + break + } + var result *analysis.TemplateCallers + if result, runErr = analysis.NewTemplateCallers(config, pkg); result != nil { + writeTo(t, &stdout, result) + } + case analysis.TemplateCallsConfiguration: + pkg, err := load.Package(dir, pl, config.TemplatesVariables) + if err != nil { + runErr = err + break + } + var result *analysis.TemplateCalls + if result, runErr = analysis.NewTemplateCalls(config, pkg); result != nil { + writeTo(t, &stdout, result) + } + default: + t.Fatalf("no analysis runs with a %T", config) + } + if stdout.Len() > 0 { + got["stdout.txt"] = relative(stdout.String()) + } + if runErr != nil { + text := runErr.Error() + if multiLine, ok := errors.AsType[muxt.MultiLineError](runErr); ok { + text = multiLine.MultiLineError() + } + got["error.txt"] = relative(text) + "\n" + } + return got +} + +func writeTo(t *testing.T, w io.Writer, result io.WriterTo) { + t.Helper() + if _, err := result.WriteTo(w); err != nil { + t.Fatal(err) + } +} + +func sortedKeys(maps ...map[string]string) []string { + var keys []string + for _, m := range maps { + for key := range m { + if !slices.Contains(keys, key) { + keys = append(keys, key) + } + } + } + slices.Sort(keys) + return keys +} diff --git a/internal/analysis/snapshots_test.go b/internal/analysis/snapshots_test.go new file mode 100644 index 00000000..00cc794e --- /dev/null +++ b/internal/analysis/snapshots_test.go @@ -0,0 +1,68 @@ +package analysis_test + +import ( + "regexp" + + "github.com/typelate/muxt/internal/analysis" +) + +// snapshots names the configuration each archive in testdata runs with, +// which also says which analysis runs: the configuration the command line +// in the comment parses into. +// +// TestCommandLineConfigurations in internal/cli states what command lines +// parse into, with literals like these; search for a literal to find its +// twin. Every configuration here is one a command line can produce: +// rejecting one that cannot work is the command line's job. +type snapshotCase struct { + archive string + config any +} + +var snapshots = []snapshotCase{ + { + // muxt list-template-callers + archive: "callers", + config: analysis.TemplateCallersConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + // muxt list-template-callers --match=^head + archive: "callers_match", + config: analysis.TemplateCallersConfiguration{TemplatesVariables: []string{"templates"}, FilterTemplates: []*regexp.Regexp{regexp.MustCompile("^head")}}, + }, + { + // muxt list-template-calls + archive: "calls", + config: analysis.TemplateCallsConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + // muxt check -v + archive: "check_bad_route_name", + config: analysis.CheckConfiguration{Verbose: true, TemplatesVariables: []string{"templates"}}, + }, + { + // muxt check + archive: "check_passes", + config: analysis.CheckConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + // muxt check + archive: "check_template_not_found", + config: analysis.CheckConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + // muxt check + archive: "check_unused_templates", + config: analysis.CheckConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + // muxt check + archive: "check_wrong_field", + config: analysis.CheckConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + // muxt --use-receiver-type=T + archive: "routes", + config: analysis.DefinitionsConfiguration{ReceiverType: "T", TemplatesVariables: []string{"templates"}}, + }, +} diff --git a/internal/analysis/template_callers.go b/internal/analysis/template_callers.go index 9194e9ce..7ec4f140 100644 --- a/internal/analysis/template_callers.go +++ b/internal/analysis/template_callers.go @@ -2,7 +2,6 @@ package analysis import ( "bytes" - "go/token" "go/types" "io" "maps" @@ -11,13 +10,16 @@ import ( "text/template/parse" "github.com/typelate/check" - - "github.com/typelate/muxt/internal/asteval" + "github.com/typelate/muxt/internal/source" ) type TemplateCallersConfiguration struct { - TemplatesVariable string - FilterTemplates []*regexp.Regexp + // TemplatesVariables are listed in order. + TemplatesVariables []string + + // FilterTemplates, when set, limits the listing to templates whose + // name matches one of them. + FilterTemplates []*regexp.Regexp } type TemplateCallers struct { @@ -34,8 +36,20 @@ func (result *TemplateCallers) WriteTo(w io.Writer) (int64, error) { } // NewTemplateCallers shows where templates are referenced -func NewTemplateCallers(config TemplateCallersConfiguration, fileSet *token.FileSet, lt *asteval.LoadedTemplates) (*TemplateCallers, error) { - global, ts := lt.Global, lt.HTML +func NewTemplateCallers(config TemplateCallersConfiguration, pkg source.Package) (*TemplateCallers, error) { + combined := &TemplateCallers{} + for _, lt := range pkg.Variables { + result, err := templateCallers(config, pkg, lt) + if err != nil { + return nil, err + } + combined.Templates = append(combined.Templates, result.Templates...) + } + return combined, nil +} + +func templateCallers(config TemplateCallersConfiguration, pkg source.Package, lt source.Variable) (*TemplateCallers, error) { + global, ts := newGlobal(pkg, lt), lt.Set refs := make(map[string][]TemplateReference) // template name -> list of references // Track {{template}} calls @@ -50,11 +64,11 @@ func NewTemplateCallers(config TemplateCallersConfiguration, fileSet *token.File } { - for c := range lt.Templates.ExecuteTemplateCalls() { - templateName, dataType := c.TemplateName, c.DataType + for _, c := range lt.Calls { + templateName, dataType := c.Template, c.Data refs[templateName] = append(refs[templateName], TemplateReference{ - Position: fileSet.Position(c.Call.Pos()), + Position: c.Position, Kind: ExecuteTemplateNode, Name: templateName, data: dataType, @@ -74,7 +88,7 @@ func NewTemplateCallers(config TemplateCallersConfiguration, fileSet *token.File if len(config.FilterTemplates) > 0 && !matchesAny(name, config.FilterTemplates) { continue } - result.Templates = append(result.Templates, NewNamedReferences(lt.Package.PkgPath, name, refs[name])) + result.Templates = append(result.Templates, NewNamedReferences(pkg.Types.Path(), name, refs[name])) } return &result, nil diff --git a/internal/analysis/template_calls.go b/internal/analysis/template_calls.go index 63f7ba55..4446258c 100644 --- a/internal/analysis/template_calls.go +++ b/internal/analysis/template_calls.go @@ -10,13 +10,16 @@ import ( "text/template/parse" "github.com/typelate/check" - - "github.com/typelate/muxt/internal/asteval" + "github.com/typelate/muxt/internal/source" ) type TemplateCallsConfiguration struct { - TemplatesVariable string - FilterTemplates []*regexp.Regexp + // TemplatesVariables are listed in order. + TemplatesVariables []string + + // FilterTemplates, when set, limits the listing to templates whose + // name matches one of them. + FilterTemplates []*regexp.Regexp } type TemplateCalls struct { @@ -33,8 +36,20 @@ func (result *TemplateCalls) WriteTo(w io.Writer) (int64, error) { } // NewTemplateCalls shows what templates use (other templates they call) -func NewTemplateCalls(config TemplateCallsConfiguration, lt *asteval.LoadedTemplates) (*TemplateCalls, error) { - global, ts := lt.Global, lt.HTML +func NewTemplateCalls(config TemplateCallsConfiguration, pkg source.Package) (*TemplateCalls, error) { + combined := &TemplateCalls{} + for _, lt := range pkg.Variables { + result, err := templateCalls(config, pkg, lt) + if err != nil { + return nil, err + } + combined.Templates = append(combined.Templates, result.Templates...) + } + return combined, nil +} + +func templateCalls(config TemplateCallsConfiguration, pkg source.Package, lt source.Variable) (*TemplateCalls, error) { + global, ts := newGlobal(pkg, lt), lt.Set // Track what each template uses (calls via {{template}}) refs := make(map[string][]TemplateReference) // template -> set of templates it calls @@ -48,10 +63,10 @@ func NewTemplateCalls(config TemplateCallsConfiguration, lt *asteval.LoadedTempl } // Analyze all templates - for c := range lt.Templates.ExecuteTemplateCalls() { - t := ts.Lookup(c.TemplateName) + for _, c := range lt.Calls { + t := ts.Lookup(c.Template) if t != nil && t.Tree != nil { - _ = check.Execute(global, t.Tree, c.DataType) + _ = check.Execute(global, t.Tree, c.Data) } } @@ -61,7 +76,7 @@ func NewTemplateCalls(config TemplateCallsConfiguration, lt *asteval.LoadedTempl if len(config.FilterTemplates) > 0 && !matchesAny(name, config.FilterTemplates) { continue } - result.Templates = append(result.Templates, NewNamedReferences(lt.Package.PkgPath, name, refs[name])) + result.Templates = append(result.Templates, NewNamedReferences(pkg.Types.Path(), name, refs[name])) } return &result, nil diff --git a/internal/analysis/testdata/callers.txtar b/internal/analysis/testdata/callers.txtar new file mode 100644 index 00000000..47d079dc --- /dev/null +++ b/internal/analysis/testdata/callers.txtar @@ -0,0 +1,48 @@ +Where each template is executed or called from, with the data type. +-- index.gohtml -- +{{define "page"}}{{template "header" .Title}}{{range .Items}}{{template "item" .}}{{end}}{{end}} +{{define "header"}}

{{.}}

{{end}} +{{define "item"}}
  • {{.Name}}
  • {{end}} +{{define "alone"}}{{template "header" "Alone"}}{{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct{ Name string } + +func Render(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} + +func RenderAlone(w io.Writer) error { + return templates.ExecuteTemplate(w, "alone", nil) +} +-- want/stdout.txt -- + +template "alone" called by: + + - server.go:17:9 execute_template "alone" untyped nil + + +template "header" called by: + + - index.gohtml:4:30 template "alone" string + - index.gohtml:1:29 template "page" string + + +template "item" called by: + + - index.gohtml:1:73 template "page" Item + + +template "page" called by: + + - server.go:13:9 execute_template "page" Page + + diff --git a/internal/analysis/testdata/callers_match.txtar b/internal/analysis/testdata/callers_match.txtar new file mode 100644 index 00000000..4ebf10fd --- /dev/null +++ b/internal/analysis/testdata/callers_match.txtar @@ -0,0 +1,33 @@ +Where the matching templates are executed or called from. +-- index.gohtml -- +{{define "page"}}{{template "header" .Title}}{{range .Items}}{{template "item" .}}{{end}}{{end}} +{{define "header"}}

    {{.}}

    {{end}} +{{define "item"}}
  • {{.Name}}
  • {{end}} +{{define "alone"}}{{template "header" "Alone"}}{{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct{ Name string } + +func Render(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} + +func RenderAlone(w io.Writer) error { + return templates.ExecuteTemplate(w, "alone", nil) +} +-- want/stdout.txt -- + +template "header" called by: + + - index.gohtml:4:30 template "alone" string + - index.gohtml:1:29 template "page" string + + diff --git a/internal/analysis/testdata/calls.txtar b/internal/analysis/testdata/calls.txtar new file mode 100644 index 00000000..bb0a12c4 --- /dev/null +++ b/internal/analysis/testdata/calls.txtar @@ -0,0 +1,38 @@ +The templates each template calls, with the data type. +-- index.gohtml -- +{{define "page"}}{{template "header" .Title}}{{range .Items}}{{template "item" .}}{{end}}{{end}} +{{define "header"}}

    {{.}}

    {{end}} +{{define "item"}}
  • {{.Name}}
  • {{end}} +{{define "alone"}}{{template "header" "Alone"}}{{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct{ Name string } + +func Render(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} + +func RenderAlone(w io.Writer) error { + return templates.ExecuteTemplate(w, "alone", nil) +} +-- want/stdout.txt -- + +template "alone" calls: + + - index.gohtml:4:30 template "header" string + + +template "page" calls: + + - index.gohtml:1:29 template "header" string + - index.gohtml:1:73 template "item" Item + + diff --git a/internal/analysis/testdata/check_bad_route_name.txtar b/internal/analysis/testdata/check_bad_route_name.txtar new file mode 100644 index 00000000..146ee14c --- /dev/null +++ b/internal/analysis/testdata/check_bad_route_name.txtar @@ -0,0 +1,34 @@ +A malformed route template name is reported with its position rather +than as an unused template. +-- index.gohtml -- +{{define "index"}}

    {{.Title}}

    {{end}} +{{define "TRACE /debug"}}{{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct { + Name string + Price int +} + +func RenderIndex(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "index", page) +} + +-- want/error.txt -- +1 error +-- want/log.txt -- + TRACE /debug + ^^^^^ +index.gohtml:2:11: TRACE method not allowed; allowed methods: GET, POST, PUT, PATCH, and DELETE + +checking endpoint index +-- want/stdout.txt -- +checked 1 diff --git a/internal/analysis/testdata/check_passes.txtar b/internal/analysis/testdata/check_passes.txtar new file mode 100644 index 00000000..990fb89f --- /dev/null +++ b/internal/analysis/testdata/check_passes.txtar @@ -0,0 +1,26 @@ +Each ExecuteTemplate call type checks the template it names, and the +templates it reaches through {{template}}. +-- index.gohtml -- +{{define "index"}}

    {{.Title}}

    {{range .Items}}{{template "item" .}}{{end}}{{end}} +{{define "item"}}
  • {{.Name}}: {{printf "%d" .Price}}
  • {{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct { + Name string + Price int +} + +func RenderIndex(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "index", page) +} + +-- want/stdout.txt -- +checked 1 diff --git a/internal/analysis/testdata/check_template_not_found.txtar b/internal/analysis/testdata/check_template_not_found.txtar new file mode 100644 index 00000000..039f6213 --- /dev/null +++ b/internal/analysis/testdata/check_template_not_found.txtar @@ -0,0 +1,30 @@ +An ExecuteTemplate call naming a template the set does not define. +-- index.gohtml -- +{{define "main"}}{{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct { + Name string + Price int +} + +func RenderIndex(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "index", page) +} + +-- want/error.txt -- +1 error +-- want/log.txt -- +server.go:16:9 ExecuteTemplate "index" Page + - template "index" not found + +-- want/stdout.txt -- +checked 1 diff --git a/internal/analysis/testdata/check_unused_templates.txtar b/internal/analysis/testdata/check_unused_templates.txtar new file mode 100644 index 00000000..9ea984e0 --- /dev/null +++ b/internal/analysis/testdata/check_unused_templates.txtar @@ -0,0 +1,35 @@ +A template nothing executes is reported; a route template is reported +as waiting for muxt generate, along with the partials only it uses. +-- index.gohtml -- +{{define "index"}}

    {{.Title}}

    {{end}} +{{define "orphan"}}

    never rendered

    {{end}} +{{define "GET /items Items()"}}{{template "row" .}}{{end}} +{{define "row"}}{{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct { + Name string + Price int +} + +func RenderIndex(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "index", page) +} + +-- want/error.txt -- +2 errors +-- want/log.txt -- +Route templates with no generated handler; run muxt generate to wire them up: + - index.gohtml:3:32: "GET /items Items()" +Unused templates: + - index.gohtml:2:20: "orphan" +-- want/stdout.txt -- +checked 1 diff --git a/internal/analysis/testdata/check_wrong_field.txtar b/internal/analysis/testdata/check_wrong_field.txtar new file mode 100644 index 00000000..bcbb0b29 --- /dev/null +++ b/internal/analysis/testdata/check_wrong_field.txtar @@ -0,0 +1,44 @@ +A field the data type does not have is reported at the call and in the +template. +-- index.gohtml -- +{{define "index"}}

    {{.Heading}}

    {{range .Items}}{{template "item" .}}{{end}}{{end}} +{{define "item"}}
  • {{.Cost}}
  • {{end}} +-- server.go -- +package server + +import "io" + +type Page struct { + Title string + Items []Item +} + +type Item struct { + Name string + Price int +} + +func RenderIndex(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "index", page) +} + +-- want/error.txt -- +1 error +-- want/log.txt -- +server.go:16:9 ExecuteTemplate "index" Page +index.gohtml:1:25: executing "index" at <.Heading>: field or method Heading not found on Page + + type Page struct { + Title string + Items []Item + } + +index.gohtml:2:24: executing "item" at <.Cost>: field or method Cost not found on Item + + type Item struct { + Name string + Price int + } + +-- want/stdout.txt -- +checked 1 diff --git a/internal/analysis/testdata/routes.txtar b/internal/analysis/testdata/routes.txtar new file mode 100644 index 00000000..d0367ffd --- /dev/null +++ b/internal/analysis/testdata/routes.txtar @@ -0,0 +1,30 @@ +The route listing shows each route template's name and source, and the +receiver's methods. +-- index.gohtml -- +{{define "GET /{$} Home()"}}

    {{.Result}}

    {{end}} +{{define "POST /save Save(form)"}}{{.Err}}{{end}} +{{define "partial"}}not a route{{end}} +-- server.go -- +package server + +import "net/url" + +type T struct{} + +func (T) Home() string { return "" } +func (*T) Save(form url.Values) error { return nil } +-- want/stdout.txt -- + +Receiver Type: example.com/server.T + +Receiver Methods: + - func (example.com/server.T) Home() string + - func (example.com/server.T) Save(form net/url.Values) error + + +Template Routes: + - POST /save Save(form) + - GET /{$} Home() + +Template Functions: + diff --git a/internal/asteval/parse.go b/internal/asteval/parse.go index acb21c4d..316907b4 100644 --- a/internal/asteval/parse.go +++ b/internal/asteval/parse.go @@ -1,9 +1,8 @@ package asteval import ( + "go/types" "text/template/parse" - - "github.com/typelate/check" ) // builtinFunctionNames are the functions text/template defines for every @@ -33,13 +32,13 @@ var builtinFunctionNames = [...]string{ // // functions supplies the names the template set registered beyond the // builtins; it may be nil. -func ParseTrees(name, text, leftDelim, rightDelim string, functions check.Functions) (map[string]*parse.Tree, error) { +func ParseTrees(name, text, leftDelim, rightDelim string, functions map[string]*types.Signature) (map[string]*parse.Tree, error) { return parse.Parse(name, text, leftDelim, rightDelim, TemplateFuncNames(functions)) } // TemplateFuncNames returns the function names a template may call, in // the shape text/template/parse wants. -func TemplateFuncNames(functions check.Functions) map[string]any { +func TemplateFuncNames(functions map[string]*types.Signature) map[string]any { names := make(map[string]any, len(functions)+len(builtinFunctionNames)) for _, name := range builtinFunctionNames { names[name] = nothing diff --git a/internal/astgen/format.go b/internal/astgen/format.go index 301dbf71..748cebe3 100644 --- a/internal/astgen/format.go +++ b/internal/astgen/format.go @@ -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 diff --git a/internal/astgen/gen.go b/internal/astgen/gen.go index 19b705d2..050ddd1c 100644 --- a/internal/astgen/gen.go +++ b/internal/astgen/gen.go @@ -19,9 +19,6 @@ type ImportManager interface { // TypeASTExpression converts a types.Type to an AST expression TypeASTExpression(tp types.Type) (ast.Expr, error) - - // Types looks up a types.Package by path - Types(pkgPath string) (*types.Package, bool) } // ExportedIdentifier creates a selector expression for an exported identifier diff --git a/internal/cli/commands.go b/internal/cli/commands.go index 1422436c..004a12c3 100644 --- a/internal/cli/commands.go +++ b/internal/cli/commands.go @@ -1,7 +1,6 @@ package cli import ( - "bytes" "cmp" _ "embed" "encoding/json" @@ -24,9 +23,10 @@ import ( "github.com/spf13/pflag" "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/fakeserver" + "github.com/typelate/muxt/internal/load" "golang.org/x/tools/go/packages" - "github.com/typelate/muxt/internal/asteval" "github.com/typelate/muxt/internal/generate" "github.com/typelate/muxt/internal/mutation" "github.com/typelate/muxt/internal/muxt" @@ -44,6 +44,34 @@ const ( ) func Commands(wd string, args []string, getEnv func(string) string, stdout, stderr io.Writer) error { + return commands(wd, args, getEnv, cliVersion, stdout, stderr, runners{ + routes: runRoutes, + check: runCheck, + generate: runGenerate, + callers: runTemplateCallers, + calls: runTemplateCalls, + mutations: runTemplateMutations, + }) +} + +// runners are what each command does once its flags have become a +// configuration. +// +// Parsing the flags, applying their defaults and rejecting what cannot +// work is the command line's job, and it is decided before a package is +// loaded. Commands runs the real runners; a test runs ones that record the +// configuration, which is how what a command line means is stated without +// loading anything. +type runners struct { + routes func(cmd *cobra.Command, wd string, config analysis.DefinitionsConfiguration) error + check func(cmd *cobra.Command, wd string, config analysis.CheckConfiguration) error + generate func(cmd *cobra.Command, wd string, config generate.RoutesFileConfiguration) error + callers func(cmd *cobra.Command, wd string, config analysis.TemplateCallersConfiguration) error + calls func(cmd *cobra.Command, wd string, config analysis.TemplateCallsConfiguration) error + mutations func(cmd *cobra.Command, wd string, config mutation.Configuration) error +} + +func commands(wd string, args []string, getEnv func(string) string, version func() (string, bool), stdout, stderr io.Writer, run runners) error { var changeDir string workingDirectory := &wd @@ -76,21 +104,7 @@ func Commands(wd string, args []string, getEnv func(string) string, stdout, stde return err } cmd.SilenceUsage = true - fileSet, pl, err := asteval.LoadPackages(*workingDirectory, rootCommandConfig.ReceiverPackage) - if err != nil { - return err - } - results, err := analysis.NewRoutes(rootCommandConfig, *workingDirectory, fileSet, pl) - if err != nil { - printMultiLineError(cmd, err) - return err - } - for _, result := range results { - if err := writeResult(cmd, cmd.OutOrStdout(), result); err != nil { - return err - } - } - return nil + return run.routes(cmd, *workingDirectory, rootCommandConfig) }, } rootCmd.PersistentFlags().StringVarP(&changeDir, "change-directory", "C", "", "change the working directory") @@ -106,14 +120,14 @@ func Commands(wd string, args []string, getEnv func(string) string, stdout, stde rootCmd.SetErr(stderr) rootCmd.AddCommand( - generateCommand(workingDirectory, getEnv), + generateCommand(workingDirectory, getEnv, version, run.generate), versionCommand(), - checkCommand(workingDirectory), - listTemplateCallersCommand(workingDirectory), - listTemplateCallsCommand(workingDirectory), + checkCommand(workingDirectory, run.check), + listTemplateCallersCommand(workingDirectory, run.callers), + listTemplateCallsCommand(workingDirectory, run.calls), exploreModuleCommand(workingDirectory), generateFakeServerCommand(workingDirectory), - testTemplateMutationsCommand(workingDirectory), + testTemplateMutationsCommand(workingDirectory, run.mutations), ) // Ensure all flag sets route their output (including deprecation warnings) to stderr @@ -127,7 +141,7 @@ func Commands(wd string, args []string, getEnv func(string) string, stdout, stde return rootCmd.Execute() } -func checkCommand(workingDirectory *string) *cobra.Command { +func checkCommand(workingDirectory *string, run func(*cobra.Command, string, analysis.CheckConfiguration) error) *cobra.Command { var ( config analysis.CheckConfiguration rt, @@ -148,25 +162,7 @@ func checkCommand(workingDirectory *string) *cobra.Command { } } cmd.SilenceUsage = true - fileSet, pl, err := asteval.LoadPackages(*workingDirectory) - if err != nil { - return err - } - logger := log.New(cmd.ErrOrStderr(), "", 0) - warnPartialAST(logger, pl) - checked, err := analysis.Check(config, *workingDirectory, logger, fileSet, pl) - if err != nil { - if printMultiLineError(cmd, err) { - return err - } - return fmt.Errorf("fail: %s", err) - } - if checked == 1 { - _, _ = fmt.Fprintln(cmd.OutOrStdout(), "ok: 1 template") - } else { - _, _ = fmt.Fprintf(cmd.OutOrStdout(), "ok: %d templates\n", checked) - } - return nil + return run(cmd, *workingDirectory, config) }, } @@ -180,7 +176,7 @@ func checkCommand(workingDirectory *string) *cobra.Command { // testTemplateMutationsCommand varies each dynamic and control flow // action in the project's templates and reports the variations the tests // let through. -func testTemplateMutationsCommand(workingDirectory *string) *cobra.Command { +func testTemplateMutationsCommand(workingDirectory *string, run func(*cobra.Command, string, mutation.Configuration) error) *cobra.Command { var ( config mutation.Configuration templatePattern string @@ -241,13 +237,7 @@ working tree is never written to.`, return err } config.SeedSet = cmd.Flags().Changed("seed") - - report, err := mutation.Run(config, *workingDirectory, cmd.ErrOrStderr()) - if err != nil { - printMultiLineError(cmd, err) - return err - } - return writeResult(cmd, cmd.OutOrStdout(), report) + return run(cmd, *workingDirectory, config) }, } @@ -293,7 +283,7 @@ func addGenerateFlagsForModule(flagSet *pflag.FlagSet, config *generate.RoutesFi // the response argument when set to a true value. const envSilenceHTTPResponseWarning = "MUXT_SILENCE_WARNING_HTTP_RESPONSE_ARGUMENT" -func generateCommand(workingDirectory *string, getEnv func(string) string) *cobra.Command { +func generateCommand(workingDirectory *string, getEnv func(string) string, version func() (string, bool), run func(*cobra.Command, string, generate.RoutesFileConfiguration) error) *cobra.Command { var ( config generate.RoutesFileConfiguration deprecatedTemplatesVar string @@ -308,7 +298,6 @@ func generateCommand(workingDirectory *string, getEnv func(string) string) *cobr return err } config.SilenceHTTPResponseWarning, _ = strconv.ParseBool(getEnv(envSilenceHTTPResponseWarning)) - stdout := cmd.OutOrStdout() for _, tv := range config.TemplatesVariables { if tv != "" && !token.IsIdentifier(tv) { return fmt.Errorf("variable %s%s", tv, errIdentSuffix) @@ -339,101 +328,12 @@ func generateCommand(workingDirectory *string, getEnv func(string) string) *cobr return fmt.Errorf("output filename must use .go extension") } - if v, ok := cliVersion(); ok && config.OutputMuxtVersion { + if v, ok := version(); ok && config.OutputMuxtVersion { config.MuxtVersion = v } applyDefaults(&config, cmd.Flags()) cmd.SilenceUsage = true - fileSet, pl, err := asteval.LoadPackages(*workingDirectory, config.ReceiverPackage) - if err != nil { - return err - } - warnPartialAST(log.New(cmd.ErrOrStderr(), "", 0), pl) - files, err := generate.TemplateRoutesFiles(*workingDirectory, config, fileSet, pl, log.New(stdout, "", 0)) - if err != nil { - printMultiLineError(cmd, err) - return err - } - - // CLEANUP HEURISTIC: - // We automatically delete muxt-generated files that are no longer needed to avoid - // manual cleanup when template files are renamed or generation modes change. - // - // Files are identified by: - // 1. Presence of "// Code generated by muxt generate" comment - // 2. Matching --output-routes-func value (to differentiate multiple route sets) - // - // Cleanup scenarios: - // - Template renamed: old_template_routes_gen.go deleted when template renamed to new.gohtml - // - Switch to single-file: all per-file *_template_routes_gen.go files deleted - // - Switch to multi-file: old single template_routes.go overwritten (if same filename) - // - Routes function unchanged: only deletes files matching current routes function - // - // IMPORTANT: If you change --output-routes-func value, old files with the previous - // routes function name will NOT be deleted (to allow multiple route sets to coexist). - // To clean up after changing routes function name, manually delete old files or - // temporarily use the old --output-routes-func value with current templates. - - // Find existing generated files for cleanup - oldGeneratedFiles, err := generate.FileArguments(*workingDirectory, config.RoutesFunction) - if err != nil { - return err - } - - for oldFilePath, oldArgs := range oldGeneratedFiles { - var ( - oldConfig generate.RoutesFileConfiguration - oldDeprecatedTemplatesVar string - ) - set := pflag.NewFlagSet("parse-old", pflag.ContinueOnError) - addGenerateFlags(set, &oldConfig, &oldDeprecatedTemplatesVar) - set.SetOutput(io.Discard) - if err := set.Parse(oldArgs); err != nil { - log.Printf("WARNING: ignored generated file %s because arguments failed to parse: %s", oldFilePath, err) - continue - } - if oldConfig.RoutesFunction != config.RoutesFunction { - delete(oldGeneratedFiles, oldFilePath) - } - if oldDeprecatedTemplatesVar != "" { - oldConfig.TemplatesVariables = []string{oldDeprecatedTemplatesVar} - } - } - - // Write new files - newGeneratedFiles := make(map[string]bool) - for i, file := range files { - var sb bytes.Buffer - writeCodeGenerationComment(&sb, configToArgs(config), config.OutputMuxtVersion) - sb.WriteString(file.Content) - if err := os.WriteFile(file.Path, sb.Bytes(), 0o644); err != nil { - for _, f := range files[:i] { - if rmErr := os.Remove(f.Path); rmErr != nil { - err = errors.Join(err, rmErr) - } - } - return err - } - // Always include the count — a uniform line parses reliably. - if file.Routes == 1 { - _, _ = fmt.Fprintf(stdout, "wrote %s: 1 route\n", filepath.Base(file.Path)) - } else { - _, _ = fmt.Fprintf(stdout, "wrote %s: %d routes\n", filepath.Base(file.Path), file.Routes) - } - newGeneratedFiles[file.Path] = true - } - - // Clean up orphaned files - // Only deletes files that match the current routes function name but weren't regenerated - for oldFile := range oldGeneratedFiles { - if !newGeneratedFiles[oldFile] { - if err := os.Remove(oldFile); err != nil && !os.IsNotExist(err) { - return fmt.Errorf("failed to remove orphaned file %s: %w", oldFile, err) - } - } - } - - return nil + return run(cmd, *workingDirectory, config) }, } @@ -513,21 +413,23 @@ func configToArgs(config generate.RoutesFileConfiguration) []string { return args } -func writeCodeGenerationComment(w io.StringWriter, args []string, includeVersion bool) { +// writeCodeGenerationComment writes the header of a generated file: the +// flags that produced it, and the version that wrote it when there is one +// to record. +func writeCodeGenerationComment(w io.StringWriter, args []string, version string) { _, _ = w.WriteString(fmt.Sprintf(codeGenerationComment, strings.TrimSpace(strings.Join(args, " ")))) - if v, ok := cliVersion(); ok && includeVersion { + if version != "" { _, _ = w.WriteString("// muxt version: ") - _, _ = w.WriteString(v) + _, _ = w.WriteString(version) _, _ = w.WriteString("\n") } // The blank line keeps the header out of the package doc comment. _, _ = w.WriteString("\n") } -func listTemplateCallersCommand(wd *string) *cobra.Command { +func listTemplateCallersCommand(wd *string, run func(*cobra.Command, string, analysis.TemplateCallersConfiguration) error) *cobra.Command { var ( config analysis.TemplateCallersConfiguration - templatesVariables []string deprecatedTemplatesVar string patterns []string ) @@ -537,7 +439,7 @@ func listTemplateCallersCommand(wd *string) *cobra.Command { Aliases: []string{"callers"}, Short: "List template callers", RunE: func(cmd *cobra.Command, args []string) error { - if err := fixTemplateVariables(&templatesVariables, deprecatedTemplatesVar); err != nil { + if err := fixTemplateVariables(&config.TemplatesVariables, deprecatedTemplatesVar); err != nil { return err } cmd.SilenceUsage = true @@ -549,38 +451,20 @@ func listTemplateCallersCommand(wd *string) *cobra.Command { config.FilterTemplates = append(config.FilterTemplates, pat) } - fileSet, pl, err := asteval.LoadPackages(*wd) - if err != nil { - return err - } - combined := &analysis.TemplateCallers{} - for _, tv := range templatesVariables { - config.TemplatesVariable = tv - lt, err := asteval.LoadTemplates(*wd, tv, pl) - if err != nil { - return err - } - result, err := analysis.NewTemplateCallers(config, fileSet, lt) - if err != nil { - return err - } - combined.Templates = append(combined.Templates, result.Templates...) - } - return writeResult(cmd, cmd.OutOrStdout(), combined) + return run(cmd, *wd, config) }, } - addUseTemplatesVarToFlagSet(cmd.Flags(), &templatesVariables, &deprecatedTemplatesVar) + addUseTemplatesVarToFlagSet(cmd.Flags(), &config.TemplatesVariables, &deprecatedTemplatesVar) cmd.Flags().StringArrayVar(&patterns, "match", nil, "filter by template name (can specify multiple regular expressions)") cmd.Flags().String("format", "text", "output format (text or json)") return cmd } -func listTemplateCallsCommand(wd *string) *cobra.Command { +func listTemplateCallsCommand(wd *string, run func(*cobra.Command, string, analysis.TemplateCallsConfiguration) error) *cobra.Command { var ( config analysis.TemplateCallsConfiguration - templatesVariables []string patterns []string deprecatedTemplatesVar string ) @@ -590,7 +474,7 @@ func listTemplateCallsCommand(wd *string) *cobra.Command { Aliases: []string{"calls"}, Short: "List template calls", RunE: func(cmd *cobra.Command, args []string) error { - if err := fixTemplateVariables(&templatesVariables, deprecatedTemplatesVar); err != nil { + if err := fixTemplateVariables(&config.TemplatesVariables, deprecatedTemplatesVar); err != nil { return err } cmd.SilenceUsage = true @@ -602,28 +486,11 @@ func listTemplateCallsCommand(wd *string) *cobra.Command { config.FilterTemplates = append(config.FilterTemplates, pat) } - _, pl, err := asteval.LoadPackages(*wd) - if err != nil { - return err - } - combined := &analysis.TemplateCalls{} - for _, tv := range templatesVariables { - config.TemplatesVariable = tv - lt, err := asteval.LoadTemplates(*wd, tv, pl) - if err != nil { - return err - } - result, err := analysis.NewTemplateCalls(config, lt) - if err != nil { - return err - } - combined.Templates = append(combined.Templates, result.Templates...) - } - return writeResult(cmd, cmd.OutOrStdout(), combined) + return run(cmd, *wd, config) }, } - addUseTemplatesVarToFlagSet(cmd.Flags(), &templatesVariables, &deprecatedTemplatesVar) + addUseTemplatesVarToFlagSet(cmd.Flags(), &config.TemplatesVariables, &deprecatedTemplatesVar) cmd.Flags().StringArrayVar(&patterns, "match", nil, "filter by template name (can specify multiple regular expressions)") cmd.Flags().String("format", "text", "output format (text or json)") @@ -673,7 +540,7 @@ func versionCommand() *cobra.Command { // the command. What it stops is the silence: without it, a package go // build rejects gets the same clean output as one that passes. func warnPartialAST(logger *log.Logger, pl []*packages.Package) { - if logger == nil || len(asteval.ParseErrors(pl)) == 0 { + if logger == nil || len(load.ParseErrors(pl)) == 0 { return } logger.Printf("warning: package has syntax errors, so these checks ran against a partial AST; run go build for the full picture") @@ -1001,12 +868,12 @@ This command is intended for exploratory use only.`, return fmt.Errorf("no muxt-generated package found at %s", dir) } - _, pl, err := asteval.LoadPackages(pkg.Dir) + _, pl, err := load.Packages(pkg.Dir) if err != nil { return err } - config := generate.FakeServerConfig{ + config := fakeserver.Config{ PackagePath: pkg.Path, PackageDir: pkg.Dir, RoutesFunction: pkg.Config.RoutesFunction, @@ -1017,7 +884,7 @@ This command is intended for exploratory use only.`, FakeImportPath: fakeImportPath, } - files, err := generate.GenerateFakeServer(config, pl) + files, err := fakeserver.Generate(config, pl) if err != nil { return err } diff --git a/internal/cli/configurations_test.go b/internal/cli/configurations_test.go new file mode 100644 index 00000000..02b3c2c3 --- /dev/null +++ b/internal/cli/configurations_test.go @@ -0,0 +1,464 @@ +package cli + +import ( + "io" + "regexp" + "strings" + "testing" + + "github.com/spf13/cobra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/generate" + "github.com/typelate/muxt/internal/mutation" +) + +// configurationOf runs a command line through the flags with runners that +// record the configuration they are handed, and returns it. Nothing is +// loaded: what a command line means is decided before a runner is called. +func configurationOf(t *testing.T, commandLine string, env map[string]string) (any, error) { + t.Helper() + var got any + record := func(config any) error { + got = config + return nil + } + err := commands("/work", strings.Fields(commandLine), func(key string) string { return env[key] }, func() (string, bool) { return "v1.2.3", true }, io.Discard, io.Discard, runners{ + routes: func(_ *cobra.Command, _ string, c analysis.DefinitionsConfiguration) error { return record(c) }, + check: func(_ *cobra.Command, _ string, c analysis.CheckConfiguration) error { return record(c) }, + generate: func(_ *cobra.Command, _ string, c generate.RoutesFileConfiguration) error { return record(c) }, + callers: func(_ *cobra.Command, _ string, c analysis.TemplateCallersConfiguration) error { return record(c) }, + calls: func(_ *cobra.Command, _ string, c analysis.TemplateCallsConfiguration) error { return record(c) }, + mutations: func(_ *cobra.Command, _ string, c mutation.Configuration) error { return record(c) }, + }) + return got, err +} + +// TestCommandLineConfigurations states what each command line parses into: +// the flags with their defaults applied and checked, as the configuration +// the command runs with. +// +// The implementation tests start from these literals. A snapshot in +// internal/generate/snapshots_test.go names the command line it stands for +// and repeats the configuration here beside it; search for the literal to +// find both. Only this test holds command lines that are rejected: a +// configuration that reaches an implementation is one that can work. +func TestCommandLineConfigurations(t *testing.T) { + for _, tt := range []struct { + name string + args string + env map[string]string + want any + }{ + { + name: "the defaults", + args: "generate", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "a receiver type", + args: "generate --use-receiver-type=T", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "a receiver type named Server", + args: "generate --use-receiver-type=Server", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "Server", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "the generated names", + args: "generate --use-receiver-type=Server --output-file=routes.go --output-routes-func=Routes --output-receiver-interface=Handlers --output-template-data-type=Data --output-template-route-paths-type=Paths", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "Routes", + ReceiverType: "Server", + ReceiverInterface: "Handlers", + TemplateDataType: "Data", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "Paths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "htmx helpers", + args: "generate --output-htmx", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputHTMX: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "the routes function parameters", + args: "generate --use-receiver-type=T --output-routes-func-with-logger-param --output-routes-func-with-path-prefix-param --output-routes-func-with-middleware-param", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + PathPrefix: true, + Logger: true, + Middleware: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "a file per template file", + args: "generate --use-receiver-type=T --output-multiple-files --output-routes-func-with-middleware-param", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + Middleware: true, + OutputMultipleFiles: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "unexported default identifiers", + args: "generate --output-exported-default-identifiers=false", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "templateRoutes", + ReceiverInterface: "routesReceiver", + TemplateDataType: "templateData", + SSETemplateDataType: "sseTemplateData", + TemplateRoutePathsTypeName: "templateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputMuxtVersion: true, + }, + }, + { + name: "unexported default identifiers keep an explicit name", + args: "generate --output-exported-default-identifiers=false --output-routes-func=Routes", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "Routes", + ReceiverInterface: "routesReceiver", + TemplateDataType: "templateData", + SSETemplateDataType: "sseTemplateData", + TemplateRoutePathsTypeName: "templateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputMuxtVersion: true, + }, + }, + { + name: "a multipart memory limit in human units", + args: "generate --use-receiver-type=T --output-multipart-max-memory=1MiB", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + MultipartMaxMemory: 1 << 20, + }, + }, + { + name: "datastar", + args: "generate --use-receiver-type=T --output-datastar", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputDatastar: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "the version left out", + args: "generate --output-muxt-version=false", + want: generate.RoutesFileConfiguration{ + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + }, + }, + { + name: "several templates variables", + args: "generate --use-templates-variable=pages --use-templates-variable=admin", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"pages", "admin"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "a receiver in another package", + args: "generate --use-receiver-type=Server --use-receiver-type-package=example.com/app", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "Server", + ReceiverPackage: "example.com/app", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "verbose", + args: "generate -v", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + Verbose: true, + }, + }, + { + name: "deprecated flags", + args: "generate --templates-variable=pages --receiver-type=T --routes-func=Routes --logger --path-prefix --output-htmx-helpers", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "Routes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"pages"}, + OutputFileName: "template_routes.go", + PathPrefix: true, + Logger: true, + OutputHTMX: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "an alias", + args: "gen --use-receiver-type=T", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + name: "the response argument warning silenced from the environment", + args: "generate --use-receiver-type=T", + env: map[string]string{envSilenceHTTPResponseWarning: "true"}, + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + SilenceHTTPResponseWarning: true, + }, + }, + { + name: "check", + args: "check", + want: analysis.CheckConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + name: "check verbosely", + args: "check -v", + want: analysis.CheckConfiguration{Verbose: true, TemplatesVariables: []string{"templates"}}, + }, + { + name: "the route listing", + args: "--use-receiver-type=T", + want: analysis.DefinitionsConfiguration{ReceiverType: "T", TemplatesVariables: []string{"templates"}}, + }, + { + name: "template callers", + args: "list-template-callers", + want: analysis.TemplateCallersConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + name: "template callers matching a name", + args: "list-template-callers --match=^head", + want: analysis.TemplateCallersConfiguration{TemplatesVariables: []string{"templates"}, FilterTemplates: []*regexp.Regexp{regexp.MustCompile("^head")}}, + }, + { + name: "template calls", + args: "list-template-calls", + want: analysis.TemplateCallsConfiguration{TemplatesVariables: []string{"templates"}}, + }, + { + name: "a mutation dry run", + args: "test-template-mutations --dry-run --seed=1 -v", + want: mutation.Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: mutation.DefaultMaxCases, Workers: 1}, + }, + { + name: "mutations with packages, patterns and go test flags", + args: "test-template-mutations --template-pattern=^page --run=TestPage --include-test-callers --max-cases=2 --workers=4 --diff=main ./... -- -count=1", + want: mutation.Configuration{TemplatesVariables: []string{"templates"}, TemplatePattern: regexp.MustCompile("^page"), Run: regexp.MustCompile("TestPage"), Packages: []string{"./..."}, GoTestArgs: []string{"-count=1"}, IncludeTests: true, MaxCases: 2, Workers: 4, Diff: "main"}, + }, + } { + t.Run(tt.name, func(t *testing.T) { + got, err := configurationOf(t, tt.args, tt.env) + require.NoError(t, err) + assert.Equal(t, tt.want, got) + }) + } +} + +// TestCommandLineRejections states the command lines that never reach a +// command: each is refused, with the error the user sees, before anything +// is loaded. +func TestCommandLineRejections(t *testing.T) { + for _, tt := range []struct { + name string + args string + wantErr string + }{ + {name: "a templates variable that is not an identifier", args: "generate --use-templates-variable=not-ok", wantErr: "variable not-ok value must be a well-formed Go identifier"}, + {name: "a routes function that is not an identifier", args: "generate --output-routes-func=1Routes", wantErr: "output-routes-func value must be a well-formed Go identifier"}, + {name: "a receiver type that is not an identifier", args: "generate --use-receiver-type=a.b", wantErr: "use-receiver-type value must be a well-formed Go identifier"}, + {name: "a receiver interface that is not an identifier", args: "generate --output-receiver-interface=a-b", wantErr: "output-receiver-interface value must be a well-formed Go identifier"}, + {name: "a template data type that is not an identifier", args: "generate --output-template-data-type=a-b", wantErr: "output-template-data-type value must be a well-formed Go identifier"}, + {name: "an sse template data type that is not an identifier", args: "generate --output-sse-template-data-type=a-b", wantErr: "output-sse-template-data-type value must be a well-formed Go identifier"}, + {name: "a route paths type that is not an identifier", args: "generate --output-template-route-paths-type=a-b", wantErr: "output-template-route-paths-type value must be a well-formed Go identifier"}, + {name: "htmx and datastar together", args: "generate --output-htmx --output-datastar", wantErr: "--output-htmx and --output-datastar are mutually exclusive; a package targets one frontend library (to mix frontends, generate separate packages that share a mux)"}, + {name: "an output file that is not Go", args: "generate --output-file=routes.txt", wantErr: "output filename must use .go extension"}, + {name: "a multipart limit of zero", args: "generate --output-multipart-max-memory=0", wantErr: `invalid argument "0" for "--output-multipart-max-memory" flag: multipart max memory must be positive, got "0"`}, + {name: "a repeated templates variable", args: "generate --use-templates-variable=pages --use-templates-variable=pages", wantErr: "duplicate template variable: pages"}, + {name: "the deprecated and new templates variable flags together", args: "check --templates-variable=a --use-templates-variable=b", wantErr: "deprecated flag templates-variable not permitted along with use-templates-variable"}, + {name: "a check templates variable that is not an identifier", args: "check --use-templates-variable=not-ok", wantErr: "variable not-ok value must be a well-formed Go identifier"}, + {name: "a template pattern that does not compile", args: "test-template-mutations --template-pattern=(", wantErr: "--template-pattern: error parsing regexp: missing closing ): `(`"}, + {name: "a run pattern that does not compile", args: "test-template-mutations --run=(", wantErr: "--run: error parsing regexp: missing closing ): `(`"}, + {name: "a callers match that does not compile", args: "list-template-callers --match=(", wantErr: "error parsing regexp: missing closing ): `(`"}, + } { + t.Run(tt.name, func(t *testing.T) { + got, err := configurationOf(t, tt.args, nil) + require.EqualError(t, err, tt.wantErr) + assert.Nil(t, got, "a rejected command line reached its command") + }) + } +} diff --git a/internal/cli/run.go b/internal/cli/run.go new file mode 100644 index 00000000..2fb6e75a --- /dev/null +++ b/internal/cli/run.go @@ -0,0 +1,214 @@ +package cli + +import ( + "bytes" + "errors" + "fmt" + "io" + "log" + "os" + "path/filepath" + + "github.com/spf13/cobra" + "github.com/spf13/pflag" + + "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/generate" + "github.com/typelate/muxt/internal/load" + "github.com/typelate/muxt/internal/mutation" +) + +// This file holds what each command does with its configuration: load the +// package, run the implementation, and write what it produced. Everything +// the flags decide happens before a runner is called, in commands.go. + +func runRoutes(cmd *cobra.Command, wd string, config analysis.DefinitionsConfiguration) error { + _, pl, err := load.Packages(wd, config.ReceiverPackage) + if err != nil { + return err + } + pkg, receiver, err := load.RoutesSource(wd, pl, config) + if err != nil { + printMultiLineError(cmd, err) + return err + } + results, err := analysis.NewRoutes(pkg, receiver) + if err != nil { + printMultiLineError(cmd, err) + return err + } + for _, result := range results { + if err := writeResult(cmd, cmd.OutOrStdout(), result); err != nil { + return err + } + } + return nil +} + +func runCheck(cmd *cobra.Command, wd string, config analysis.CheckConfiguration) error { + _, pl, err := load.Packages(wd) + if err != nil { + return err + } + logger := log.New(cmd.ErrOrStderr(), "", 0) + warnPartialAST(logger, pl) + checked, err := func() (int, error) { + pkg, err := load.Package(wd, pl, config.TemplatesVariables) + if err != nil { + return 0, err + } + return analysis.Check(config, logger, pkg) + }() + if err != nil { + if printMultiLineError(cmd, err) { + return err + } + return fmt.Errorf("fail: %s", err) + } + if checked == 1 { + _, _ = fmt.Fprintln(cmd.OutOrStdout(), "ok: 1 template") + } else { + _, _ = fmt.Fprintf(cmd.OutOrStdout(), "ok: %d templates\n", checked) + } + return nil +} + +func runTemplateCallers(cmd *cobra.Command, wd string, config analysis.TemplateCallersConfiguration) error { + _, pl, err := load.Packages(wd) + if err != nil { + return err + } + pkg, err := load.Package(wd, pl, config.TemplatesVariables) + if err != nil { + return err + } + result, err := analysis.NewTemplateCallers(config, pkg) + if err != nil { + return err + } + return writeResult(cmd, cmd.OutOrStdout(), result) +} + +func runTemplateCalls(cmd *cobra.Command, wd string, config analysis.TemplateCallsConfiguration) error { + _, pl, err := load.Packages(wd) + if err != nil { + return err + } + pkg, err := load.Package(wd, pl, config.TemplatesVariables) + if err != nil { + return err + } + result, err := analysis.NewTemplateCalls(config, pkg) + if err != nil { + return err + } + return writeResult(cmd, cmd.OutOrStdout(), result) +} + +func runTemplateMutations(cmd *cobra.Command, wd string, config mutation.Configuration) error { + report, err := mutation.Run(config, wd, cmd.ErrOrStderr()) + if err != nil { + printMultiLineError(cmd, err) + return err + } + return writeResult(cmd, cmd.OutOrStdout(), report) +} + +func runGenerate(cmd *cobra.Command, wd string, config generate.RoutesFileConfiguration) error { + stdout := cmd.OutOrStdout() + _, pl, err := load.Packages(wd, config.ReceiverPackage) + if err != nil { + return err + } + warnPartialAST(log.New(cmd.ErrOrStderr(), "", 0), pl) + pkg, receiver, err := load.GenerateSource(wd, pl, config) + if err != nil { + printMultiLineError(cmd, err) + return err + } + files, err := generate.TemplateRoutesFiles(wd, config, pkg, receiver, log.New(stdout, "", 0)) + if err != nil { + printMultiLineError(cmd, err) + return err + } + + // CLEANUP HEURISTIC: + // We automatically delete muxt-generated files that are no longer needed to avoid + // manual cleanup when template files are renamed or generation modes change. + // + // Files are identified by: + // 1. Presence of "// Code generated by muxt generate" comment + // 2. Matching --output-routes-func value (to differentiate multiple route sets) + // + // Cleanup scenarios: + // - Template renamed: old_template_routes_gen.go deleted when template renamed to new.gohtml + // - Switch to single-file: all per-file *_template_routes_gen.go files deleted + // - Switch to multi-file: old single template_routes.go overwritten (if same filename) + // - Routes function unchanged: only deletes files matching current routes function + // + // IMPORTANT: If you change --output-routes-func value, old files with the previous + // routes function name will NOT be deleted (to allow multiple route sets to coexist). + // To clean up after changing routes function name, manually delete old files or + // temporarily use the old --output-routes-func value with current templates. + + // Find existing generated files for cleanup + oldGeneratedFiles, err := generate.FileArguments(wd, config.RoutesFunction) + if err != nil { + return err + } + + for oldFilePath, oldArgs := range oldGeneratedFiles { + var ( + oldConfig generate.RoutesFileConfiguration + oldDeprecatedTemplatesVar string + ) + set := pflag.NewFlagSet("parse-old", pflag.ContinueOnError) + addGenerateFlags(set, &oldConfig, &oldDeprecatedTemplatesVar) + set.SetOutput(io.Discard) + if err := set.Parse(oldArgs); err != nil { + log.Printf("WARNING: ignored generated file %s because arguments failed to parse: %s", oldFilePath, err) + continue + } + if oldConfig.RoutesFunction != config.RoutesFunction { + delete(oldGeneratedFiles, oldFilePath) + } + if oldDeprecatedTemplatesVar != "" { + oldConfig.TemplatesVariables = []string{oldDeprecatedTemplatesVar} + } + } + + // Write new files + newGeneratedFiles := make(map[string]bool) + for i, file := range files { + var sb bytes.Buffer + writeCodeGenerationComment(&sb, configToArgs(config), config.MuxtVersion) + sb.WriteString(file.Content) + if err := os.WriteFile(file.Path, sb.Bytes(), 0o644); err != nil { + for _, f := range files[:i] { + if rmErr := os.Remove(f.Path); rmErr != nil { + err = errors.Join(err, rmErr) + } + } + return err + } + // Always include the count — a uniform line parses reliably. + if file.Routes == 1 { + _, _ = fmt.Fprintf(stdout, "wrote %s: 1 route\n", filepath.Base(file.Path)) + } else { + _, _ = fmt.Fprintf(stdout, "wrote %s: %d routes\n", filepath.Base(file.Path), file.Routes) + } + newGeneratedFiles[file.Path] = true + } + + // Clean up orphaned files + // Only deletes files that match the current routes function name but weren't regenerated + for oldFile := range oldGeneratedFiles { + if !newGeneratedFiles[oldFile] { + if err := os.Remove(oldFile); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("failed to remove orphaned file %s: %w", oldFile, err) + } + } + } + + return nil +} diff --git a/internal/generate/fake_server.go b/internal/fakeserver/fakeserver.go similarity index 86% rename from internal/generate/fake_server.go rename to internal/fakeserver/fakeserver.go index 3b4d98fe..0aaf0ac6 100644 --- a/internal/generate/fake_server.go +++ b/internal/fakeserver/fakeserver.go @@ -1,17 +1,17 @@ -package generate +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" ) -// FakeServerConfig holds the configuration for generating a fake server. -type FakeServerConfig struct { +// Config holds the configuration for generating a fake server. +type Config struct { PackagePath string // import path of the muxt-generated package PackageDir string // absolute directory of the package RoutesFunction string // e.g. "TemplateRoutes" @@ -22,20 +22,20 @@ type FakeServerConfig struct { FakeImportPath string // import path of the generated fake package } -// FakeServerFiles holds the generated files for the fake server. -type FakeServerFiles struct { +// Files holds the generated files for the fake server. +type Files struct { Main []byte // main.go — small, readable entry point Fake []byte // internal/fake/receiver.go — counterfeiter-generated fake struct } -// GenerateFakeServer generates two files: a main.go with the httptest server +// Generate generates two files: a main.go with the httptest server // entry point, and a receiver.go with the counterfeiter-generated fake struct // in a separate package. // // The target package must not be "main" — the fake server imports it as a library. // // The fake implementation interface is unstable and should not be relied upon. -func GenerateFakeServer(config FakeServerConfig, pl []*packages.Package) (*FakeServerFiles, error) { +func Generate(config Config, pl []*packages.Package) (*Files, error) { // Find the target package and validate it's not main. var targetPkg *packages.Package for _, pkg := range pl { @@ -82,22 +82,22 @@ func GenerateFakeServer(config FakeServerConfig, pl []*packages.Package) (*FakeS } type mainTemplateData struct { - FakeServerConfig + Config PackageName string } var mainBuf bytes.Buffer if err := mainFuncTemplate.Execute(&mainBuf, mainTemplateData{ - FakeServerConfig: config, - PackageName: pkgAlias, + Config: config, + PackageName: pkgAlias, }); 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 &FakeServerFiles{ + return &Files{ Main: mainBytes, Fake: fakeSource, }, nil diff --git a/internal/generate/file.go b/internal/generate/file.go index 1d082ebf..134dff69 100644 --- a/internal/generate/file.go +++ b/internal/generate/file.go @@ -3,7 +3,6 @@ package generate import ( "crypto/sha1" "encoding/hex" - "fmt" "go/ast" "go/parser" "go/token" @@ -11,74 +10,35 @@ import ( "log" "maps" "path" - "path/filepath" "slices" "strconv" "strings" - "golang.org/x/tools/go/packages" - - "github.com/typelate/muxt/internal/asteval" + "github.com/typelate/muxt/internal/source" ) +// File is one generated Go file: the package it is written into, which the +// types it names are qualified 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 - files map[string]*ast.File - packages []*packages.Package - outPkg *packages.Package + pkg source.Package packageIdentifiers map[string]string importSpecs []*ast.ImportSpec } -func newFile(filePath string, fileSet *token.FileSet, list []*packages.Package) (*File, error) { - if fileSet == nil { - fileSet = token.NewFileSet() - } - file := &File{ - fileSet: fileSet, - typesCache: make(map[string]*types.Package), - files: make(map[string]*ast.File), - packages: make([]*packages.Package, 0), +func newFile(pkg source.Package) *File { + return &File{ + pkg: pkg, packageIdentifiers: make(map[string]string), } - file.addPackages(list) - pkg, found := asteval.PackageAtFilepath(list, filePath) - if !found { - // filePath names the output file, which need not exist yet; the - // lookup is for the package in its directory. - return nil, asteval.NoPackageError(filepath.Dir(filePath), list) - } - file.outPkg = pkg - return file, nil -} - -func (file *File) Package(path string) (*packages.Package, bool) { - return asteval.PackageWithPath(file.packages, path) -} - -func (file *File) addPackages(packages []*packages.Package) { - file.packages = slices.Grow(file.packages, len(packages)) - for _, pkg := range packages { - if pkg == nil { - continue - } - file.typesCache[pkg.PkgPath] = pkg.Types - file.packages = append(file.packages, pkg) - } } -func (file *File) OutputPackage() *packages.Package { return file.outPkg } - -// Packages returns the loaded package list backing this file's type lookups. -func (file *File) Packages() []*packages.Package { return file.packages } - -func (file *File) SyntaxFile(pos token.Pos) (*ast.File, *token.FileSet, error) { - position := file.fileSet.Position(pos) - fSet := token.NewFileSet() - f, err := parser.ParseFile(fSet, position.Filename, nil, parser.AllErrors|parser.ParseComments|parser.SkipObjectResolution) - return f, fSet, err -} +// OutputPackage is the package the generated file is written into. +func (file *File) OutputPackage() source.Package { return file.pkg } func (file *File) TypeASTExpression(tp types.Type) (ast.Expr, error) { s := types.TypeString(tp, file.pkgQualifier) @@ -87,67 +47,14 @@ func (file *File) TypeASTExpression(tp types.Type) (ast.Expr, error) { // pkgQualifier implements types.Qualifier func (file *File) pkgQualifier(pkg *types.Package) string { - if pkg.Path() == file.outPkg.PkgPath { + if pkg.Path() == file.pkg.Types.Path() { return "" } return file.Import(pkg.Name(), pkg.Path()) } -func (file *File) StructField(pos token.Pos) (*ast.Field, error) { - f, fileSet, err := file.SyntaxFile(pos) - if err != nil { - return nil, err - } - position := file.fileSet.Position(pos) - for _, d := range f.Decls { - switch decl := d.(type) { - case *ast.GenDecl: - for _, s := range decl.Specs { - switch spec := s.(type) { - case *ast.TypeSpec: - tp, ok := spec.Type.(*ast.StructType) - if !ok { - continue - } - - for _, field := range tp.Fields.List { - for _, name := range field.Names { - p := fileSet.Position(name.Pos()) - if p != position { - continue - } - return field, nil - } - } - } - } - } - } - return nil, fmt.Errorf("failed to find field") -} - -func (file *File) Types(pkgPath string) (*types.Package, bool) { - if p, ok := file.typesCache[pkgPath]; ok { - return p, true - } - for _, pkg := range file.packages { - if pkg.Types.Path() == pkgPath { - p := pkg.Types - file.typesCache[pkgPath] = p - return p, true - } - } - for _, pkg := range file.packages { - if p, ok := recursivelySearchImports(pkg.Types, pkgPath); ok { - file.typesCache[pkgPath] = p - return p, true - } - } - return nil, false -} - func (file *File) Import(pkgIdent, pkgPath string) string { - if pkgPath == file.outPkg.PkgPath { + if pkgPath == file.pkg.Types.Path() { log.Fatal("package path cannot be the same as the output package") return "" } @@ -160,20 +67,6 @@ func (file *File) ImportSpecs() []*ast.ImportSpec { return slices.CompactFunc(result, func(a, b *ast.ImportSpec) bool { return a.Path.Value == b.Path.Value }) } -func recursivelySearchImports(pt *types.Package, pkgPath string) (*types.Package, bool) { - for _, pkg := range pt.Imports() { - if pkg.Path() == pkgPath { - return pkg, true - } - } - for _, pkg := range pt.Imports() { - if im, ok := recursivelySearchImports(pkg, pkgPath); ok { - return im, true - } - } - return nil, false -} - func packageImportName(importSpecs *[]*ast.ImportSpec, packageIdentifiers map[string]string, pkgPath, pkgIdent string) string { if ident, ok := packageIdentifiers[pkgPath]; ok { return ident diff --git a/internal/generate/file_test.go b/internal/generate/file_test.go index 9463eff3..5f3ef353 100644 --- a/internal/generate/file_test.go +++ b/internal/generate/file_test.go @@ -3,40 +3,20 @@ package generate import ( "go/ast" "go/token" - "os" - "path/filepath" - "sync" + "go/types" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "golang.org/x/tools/go/packages" "github.com/typelate/muxt/internal/astgen" + "github.com/typelate/muxt/internal/source" ) -var ( - workingDir = sync.OnceValues(func() (string, error) { - return os.Getwd() - }) - fileSet = sync.OnceValue(func() *token.FileSet { - return token.NewFileSet() - }) - loadPkg = sync.OnceValues(func() ([]*packages.Package, error) { - wd, err := workingDir() - if err != nil { - return nil, err - } - return loadPackages(wd, []string{"context", "net/http", wd}) - }) -) - -func loadPackages(wd string, patterns []string) ([]*packages.Package, error) { - return packages.Load(&packages.Config{ - Fset: fileSet(), - Mode: packages.NeedModule | packages.NeedName | packages.NeedFiles | packages.NeedTypes | packages.NeedSyntax | packages.NeedEmbedPatterns | packages.NeedEmbedFiles, - Dir: wd, - }, patterns...) +// outputFile is a File for a package that imports nothing: what the +// import bookkeeping needs, and no more. +func outputFile() *File { + return newFile(source.Package{Types: types.NewPackage("example.com/server", "server")}) } func TestImports(t *testing.T) { @@ -48,54 +28,22 @@ func TestImports(t *testing.T) { return astgen.Format(decl) } t.Run("initial add", func(t *testing.T) { - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() assert.Equal(t, "http", file.Import("http", "net/http")) assert.Equal(t, genDecl(file), `import "net/http"`) }) t.Run("initial with pkg ident", func(t *testing.T) { - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() assert.Equal(t, "p", file.Import("p", "net/http")) assert.Equal(t, genDecl(file), `import p "net/http"`) }) t.Run("initial with empty ident", func(t *testing.T) { - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() assert.Equal(t, "http", file.Import("", "net/http")) assert.Equal(t, genDecl(file), `import "net/http"`) }) t.Run("initial with empty ident", func(t *testing.T) { - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() _ = file.Import("", "net/http") _ = file.Import("", "html/template") assert.Equal(t, genDecl(file), `import ( @@ -104,15 +52,7 @@ func TestImports(t *testing.T) { )`) }) t.Run("it respects order", func(t *testing.T) { - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() _ = file.Import("", "html/template") _ = file.Import("", "net/http") assert.Equal(t, genDecl(file), `import ( @@ -121,42 +61,19 @@ func TestImports(t *testing.T) { )`) }) t.Run("it returns the registered identifier", func(t *testing.T) { - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() _ = file.Import("t", "html/template") assert.Equal(t, "t", file.Import("", "html/template")) }) t.Run("it returns the package path base", func(t *testing.T) { - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() _ = file.Import("", "html/template") assert.Equal(t, "template", file.Import("", "html/template")) }) } func TestHTTPStatusCode(t *testing.T) { - fSet := fileSet() - wd, err := workingDir() - require.NoError(t, err) - pl, err := loadPackages(wd, []string{wd}) - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() exp := astgen.HTTPStatusCode(file, 600) require.NotNil(t, exp) diff --git a/internal/generate/format.go b/internal/generate/format.go new file mode 100644 index 00000000..eec1fdb1 --- /dev/null +++ b/internal/generate/format.go @@ -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 +} diff --git a/internal/generate/format_test.go b/internal/generate/format_test.go new file mode 100644 index 00000000..fb29bbca --- /dev/null +++ b/internal/generate/format_test.go @@ -0,0 +1,72 @@ +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"), + }, + want: "package server\n\nimport (\n\t\"bytes\"\n\t\"net/http\"\n\t\"server/internal/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) + } + }) + } +} diff --git a/internal/generate/generated_test.go b/internal/generate/generated_test.go new file mode 100644 index 00000000..efef5384 --- /dev/null +++ b/internal/generate/generated_test.go @@ -0,0 +1,38 @@ +package generate + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestFileArguments(t *testing.T) { + dir := t.TempDir() + for name, content := range map[string]string{ + "template_routes.go": "// Code generated by muxt generate --use-receiver-type=T. DO NOT EDIT.\n\npackage server\n", + "index_gen.go": "// Code generated by muxt generate . DO NOT EDIT.\n", + "other_generated.go": "// Code generated by stringer. DO NOT EDIT.\n", + "handwritten.go": "package server\n", + "notes.txt": "// Code generated by muxt generate --output-file=notes.txt. DO NOT EDIT.\n", + "second_line.go": "package server\n// Code generated by muxt generate --output-htmx. DO NOT EDIT.\n", + "empty.go": "", + "sub/template_rou.go": "// Code generated by muxt generate --output-htmx. DO NOT EDIT.\n", + } { + path := filepath.Join(dir, filepath.FromSlash(name)) + require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755)) + require.NoError(t, os.WriteFile(path, []byte(content), 0o644)) + } + + files, err := FileArguments(dir, DefaultRoutesFunctionName) + require.NoError(t, err) + assert.Equal(t, map[string][]string{ + filepath.Join(dir, "template_routes.go"): {"--use-receiver-type=T"}, + filepath.Join(dir, "index_gen.go"): {}, + }, files, "only the .go files in the directory itself whose first line is the muxt header") + + _, err = FileArguments(filepath.Join(dir, "missing"), DefaultRoutesFunctionName) + assert.Error(t, err) +} diff --git a/internal/generate/groups.go b/internal/generate/groups.go index 09c47a45..fbb5676f 100644 --- a/internal/generate/groups.go +++ b/internal/generate/groups.go @@ -5,10 +5,8 @@ import ( "path/filepath" "strings" - "golang.org/x/tools/go/packages" - - "github.com/typelate/muxt/internal/asteval" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) type templateGroups struct { @@ -17,17 +15,12 @@ type templateGroups struct { all []muxt.Definition } -func groupTemplates(wd string, config RoutesFileConfiguration, routesPkg *packages.Package) (templateGroups, error) { +func groupTemplates(config RoutesFileConfiguration, variables []source.Variable) (templateGroups, error) { result := templateGroups{ byFile: make(map[string][]muxt.Definition), } - for _, tv := range config.TemplatesVariables { - lt, ts, err := asteval.HTMLTemplates(tv, routesPkg) - if err != nil { - return result, err - } - - defs, err := muxt.Definitions(ts, tv, lt) + for _, tv := range variables { + defs, err := muxt.Definitions(tv) if err != nil { return result, err } diff --git a/internal/generate/html.go b/internal/generate/html.go index 3ba7c668..ae22faf5 100644 --- a/internal/generate/html.go +++ b/internal/generate/html.go @@ -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 @@ -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())) @@ -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 diff --git a/internal/generate/routes.go b/internal/generate/routes.go index 7d253cf5..5653dba9 100644 --- a/internal/generate/routes.go +++ b/internal/generate/routes.go @@ -15,11 +15,11 @@ import ( "strings" "github.com/ettle/strcase" - "golang.org/x/tools/go/packages" "github.com/typelate/muxt/internal/asteval" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) const ( @@ -97,33 +97,26 @@ type RoutesFileConfiguration struct { // request.ParseMultipartForm when no override is set. const DefaultMultipartMaxMemory int64 = 32 << 20 -func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, fileSet *token.FileSet, pl []*packages.Package, logger *log.Logger) ([]GeneratedFile, error) { +// TemplateRoutesFiles generates the routes files for pkg, written into wd: +// the package the files belong to, which is the one in the output file's +// directory. receiver is the type --use-receiver-type named, or nil when +// handler methods are inferred from the templates. +func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, pkg source.Package, receiver *types.Named, logger *log.Logger) ([]GeneratedFile, error) { if !token.IsIdentifier(config.PackageName) { return nil, fmt.Errorf("package name %q is not an identifier", config.PackageName) } - file, err := newFile(filepath.Join(wd, config.OutputFileName), fileSet, pl) - if err != nil { - return nil, err - } - routesPkg := file.OutputPackage() + file := newFile(pkg) - config.PackagePath = routesPkg.PkgPath - config.PackageName = routesPkg.Name + config.PackagePath = pkg.Types.Path() + config.PackageName = pkg.Types.Name() config.SSETemplateDataType = cmp.Or(config.SSETemplateDataType, "SSETemplateData") - var receiver *types.Named - if config.ReceiverType == "" { - receiver = asteval.NamedEmptyStruct("Receiver", routesPkg.Types) - } else { - receiverPkgPath := cmp.Or(config.ReceiverPackage, config.PackagePath) - receiver, err = asteval.FindType(pl, receiverPkgPath, config.ReceiverType) - if err != nil { - return nil, err - } + if receiver == nil { + receiver = asteval.NamedEmptyStruct("Receiver", pkg.Types) } - groups, err := groupTemplates(wd, config, routesPkg) + groups, err := groupTemplates(config, pkg.Variables) if err != nil { return nil, err } @@ -185,7 +178,7 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, fileSet *tok generatedFiles []GeneratedFile ) if config.OutputMultipleFiles { - files, err := sourceFileRouteFunctionFiles(wd, config, templateSourceFiles, groups, logger, file, receiver, routesPkg, receiverInterface, routesFunc) + files, err := sourceFileRouteFunctionFiles(wd, config, templateSourceFiles, groups, logger, file, receiver, receiverInterface, routesFunc) if err != nil { return files, err } @@ -200,7 +193,7 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, fileSet *tok } // Generate handlers for parse-based templates (empty sourceFile) - if err := hydrateGroup(topLevelTemplateRoutes, file, receiver, routesPkg.Types, receiverInterface, logger, config.ReceiverType != "" && logger != nil, logger != nil && !config.SilenceHTTPResponseWarning); err != nil { + if err := hydrateGroup(topLevelTemplateRoutes, file, receiver, receiverInterface, logger, config.ReceiverType != "" && logger != nil, logger != nil && !config.SilenceHTTPResponseWarning); err != nil { return nil, err } for _, def := range topLevelTemplateRoutes { @@ -236,17 +229,11 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, fileSet *tok }, }) - is := file.ImportSpecs() - importSpecs := make([]ast.Spec, 0, len(is)) - for _, s := range is { - importSpecs = append(importSpecs, s) - } + // The import declaration is filled in last: building the other + // declarations is what registers the imports they use. + importDecl := &ast.GenDecl{Tok: token.IMPORT} decls := []ast.Decl{ - // import - &ast.GenDecl{ - Tok: token.IMPORT, - Specs: importSpecs, - }, + importDecl, // type &ast.GenDecl{ @@ -268,13 +255,16 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, fileSet *tok decls = append(decls, sseTemplateDataDecls(file, config)...) } decls = append(decls, routePathDecls...) + for _, spec := range file.ImportSpecs() { + importDecl.Specs = append(importDecl.Specs, spec) + } outputFile := &ast.File{ Name: ast.NewIdent(config.PackageName), Decls: decls, } filePath := filepath.Join(wd, config.OutputFileName) - content, err := astgen.FormatFile(filePath, outputFile) + content, err := formatFile(filePath, outputFile) if err != nil { return nil, err } @@ -292,7 +282,7 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, fileSet *tok // signatures — one line per method and the explanation once after the // list; the default mode synthesizes every method by design, so it // stays quiet. -func hydrateGroup(defs []muxt.Definition, file *File, receiver *types.Named, templatesPackage *types.Package, receiverInterface *ast.InterfaceType, logger *log.Logger, noteSynthesized, warnResponse bool) error { +func hydrateGroup(defs []muxt.Definition, file *File, receiver *types.Named, receiverInterface *ast.InterfaceType, logger *log.Logger, noteSynthesized, warnResponse bool) error { var resolveErrs []error synthesized := 0 for i := range defs { @@ -304,7 +294,7 @@ func hydrateGroup(defs []muxt.Definition, file *File, receiver *types.Named, tem if defs[i].FunctionIdentifier() == nil { continue } - if err := muxt.ResolveCall(&defs[i], templatesPackage, receiver, file.Packages()); err != nil { + if err := muxt.ResolveCall(&defs[i], file.OutputPackage(), receiver); err != nil { resolveErrs = append(resolveErrs, err) continue } @@ -357,7 +347,7 @@ func accumulateReceiverMethods(name string, sig *types.Signature, isMethod bool, return nil } -func sourceFileRouteFunctionFiles(wd string, config RoutesFileConfiguration, templateSourceFiles []string, groups templateGroups, logger *log.Logger, file *File, receiver *types.Named, routesPkg *packages.Package, receiverInterface *ast.InterfaceType, routesFunc *ast.FuncDecl) ([]GeneratedFile, error) { +func sourceFileRouteFunctionFiles(wd string, config RoutesFileConfiguration, templateSourceFiles []string, groups templateGroups, logger *log.Logger, file *File, receiver *types.Named, receiverInterface *ast.InterfaceType, routesFunc *ast.FuncDecl) ([]GeneratedFile, error) { var generatedFiles []GeneratedFile for _, sourceFile := range templateSourceFiles { definitions := groups.byFile[sourceFile] @@ -369,7 +359,7 @@ func sourceFileRouteFunctionFiles(wd string, config RoutesFileConfiguration, tem receiverInterfaceName := strcase.ToGoCamel(fileIdentifier + " " + config.ReceiverInterface) routesFuncName := strcase.ToGoCamel(fileIdentifier + " " + config.RoutesFunction) - perFileAST, err := generatePerFileAST(sourceFile, definitions, file, routesFuncName, receiverInterfaceName, logger, config, receiver, routesPkg) + perFileAST, err := generatePerFileAST(sourceFile, definitions, newFile(file.OutputPackage()), routesFuncName, receiverInterfaceName, logger, config, receiver) if err != nil { return nil, fmt.Errorf("failed to generate routes for %s: %w", sourceFile, err) } @@ -380,7 +370,7 @@ func sourceFileRouteFunctionFiles(wd string, config RoutesFileConfiguration, tem outputFileName := baseFileName + "_template_routes_gen.go" outputFilePath := filepath.Join(wd, outputFileName) - content, err := astgen.FormatFile(outputFilePath, perFileAST) + content, err := formatFile(outputFilePath, perFileAST) if err != nil { return nil, fmt.Errorf("failed to format %s: %w", outputFileName, err) } @@ -491,7 +481,6 @@ func generatePerFileRouteFunction( config RoutesFileConfiguration, receiver *types.Named, receiverInterface *ast.InterfaceType, - routesPkg *packages.Package, ) (*ast.FuncDecl, error) { if sourceFile == "" { return nil, fmt.Errorf("sourceFile cannot be empty") @@ -536,7 +525,7 @@ func generatePerFileRouteFunction( } // Generate handlers for each template - if err := hydrateGroup(defs, file, receiver, routesPkg.Types, receiverInterface, logger, config.ReceiverType != "" && logger != nil, logger != nil && !config.SilenceHTTPResponseWarning); err != nil { + if err := hydrateGroup(defs, file, receiver, receiverInterface, logger, config.ReceiverType != "" && logger != nil, logger != nil && !config.SilenceHTTPResponseWarning); err != nil { return nil, err } for i := range defs { @@ -571,7 +560,6 @@ func generatePerFileAST( logger *log.Logger, config RoutesFileConfiguration, receiver *types.Named, - routesPkg *packages.Package, ) (*ast.File, error) { if sourceFile == "" { return nil, fmt.Errorf("sourceFile cannot be empty") @@ -592,7 +580,6 @@ func generatePerFileAST( config, receiver, scopedReceiverInterface, - routesPkg, ) if err != nil { return nil, err @@ -760,6 +747,27 @@ func callWriteOnResponse(bufferIdent string) *ast.AssignStmt { } } +// cloneCall copies a template name's call expression deeply enough that +// rewriting the copy's arguments, or a nested call's function, leaves the +// original untouched. A template name's call has only identifiers and +// nested calls for arguments. +func cloneCall(call *ast.CallExpr) *ast.CallExpr { + clone := *call + clone.Args = make([]ast.Expr, len(call.Args)) + for i, arg := range call.Args { + switch arg := arg.(type) { + case *ast.CallExpr: + clone.Args[i] = cloneCall(arg) + case *ast.Ident: + ident := *arg + clone.Args[i] = &ident + default: + clone.Args[i] = arg + } + } + return &clone +} + func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, file *File, resultType types.Type, signature *types.Signature, args []muxt.Argument, parsed map[string]struct{}, rdIdent string, config RoutesFileConfiguration, call *ast.CallExpr, validationFailureBlock ValidationErrorBlock, parseErrBlock func() *ast.BlockStmt) ([]ast.Stmt, error) { if parseErrBlock == nil { // Normal handlers accumulate scalar-parse failures into the template @@ -841,7 +849,7 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f continue } name := arg.Name - argType, ok := muxt.DefaultScopeType(file.Packages(), &def, name) + argType, ok := muxt.DefaultScopeType(file.OutputPackage(), &def, name) if !ok { return nil, fmt.Errorf("failed to determine type for %s", name) } @@ -890,7 +898,6 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f return nil, err } statements = append(statements, s...) - def.SetArgumentType(name, param.Type()) case name == muxt.TemplateNameScopeIdentifierLastEventID: parsed[name] = struct{}{} s, err := generateParseValueFromStringStatements(file, def, name+"Parsed", resultType, src, param.Type(), nil, singleAssignment(token.DEFINE, ast.NewIdent(ident)), parseErrBlock()) @@ -898,7 +905,6 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f return nil, err } statements = append(statements, s...) - def.SetArgumentType(name, param.Type()) case arg.Name == muxt.TemplateNameScopeIdentifierForm: s, err := appendParseFormToStructStatements(statements, def, file, resultType, arg, args[i], validationFailureBlock, parseErrBlock) if err != nil { @@ -1147,7 +1153,7 @@ func generateParseValueFromStringStatements(file *File, _ muxt.Definition, tmp s Args: []ast.Expr{exp}, }) } - switch muxt.UnmarshalMethodFor(file.Packages(), valueType) { + switch muxt.UnmarshalMethodFor(file.OutputPackage(), valueType) { case muxt.UnmarshalBool: return parseBlock(tmp, astgen.StrconvParseBoolCall(file, str), validations, errBlock, assignment), nil case muxt.UnmarshalInt: diff --git a/internal/generate/routes_test.go b/internal/generate/routes_test.go new file mode 100644 index 00000000..59a17260 --- /dev/null +++ b/internal/generate/routes_test.go @@ -0,0 +1,50 @@ +package generate + +import ( + "testing" + + "github.com/typelate/muxt/internal/astgen" + "github.com/typelate/muxt/internal/muxt" +) + +// TestHandlerGenerationLeavesTheRouteAsResolved generates one handler +// twice. Rendering argument parsing rewrites the call to the locals it +// declares; were that done to the route itself, the second handler would +// be generated from the first one's rewrites. +func TestHandlerGenerationLeavesTheRouteAsResolved(t *testing.T) { + config := testConfig() + config.ReceiverType = "T" + pkg, receiver := testSource(t, `package server + +type T struct{} + +func (T) Article(id int, title string) (string, error) { return "", nil } + +func (T) Title(id int) string { return "" } +`, "T", `{{define "GET /article/{id} Article(id, Title(id))"}}{{end}}`) + + groups, err := groupTemplates(config, pkg.Variables) + if err != nil { + t.Fatal(err) + } + file := newFile(pkg) + def := groups.all[0] + if err := muxt.ResolveCall(&def, pkg, receiver); err != nil { + t.Fatal(err) + } + + var handlers []string + for range 2 { + handler, err := callHandlerFunc(file, config, def, config.ReceiverInterface) + if err != nil { + t.Fatal(err) + } + handlers = append(handlers, astgen.Format(handler)) + } + if handlers[0] != handlers[1] { + t.Errorf("the second handler differs from the first:\n%s\nsecond:\n%s", handlers[0], handlers[1]) + } + if got := astgen.Format(def.CallExpression()); got != "Article(id, Title(id))" { + t.Errorf("the route's call is %s after generation, want it as written", got) + } +} diff --git a/internal/generate/snapshot_test.go b/internal/generate/snapshot_test.go new file mode 100644 index 00000000..8c96e7f3 --- /dev/null +++ b/internal/generate/snapshot_test.go @@ -0,0 +1,213 @@ +package generate_test + +import ( + "flag" + "go/ast" + "go/parser" + "go/token" + "log" + "os" + "path" + "path/filepath" + "slices" + "strconv" + "strings" + "testing" + + "golang.org/x/tools/txtar" + + "github.com/typelate/muxt/internal/generate" + "github.com/typelate/muxt/internal/load" + "github.com/typelate/muxt/internal/load/loadtest" + "github.com/typelate/muxt/internal/muxt" +) + +var update = flag.Bool("update", false, "rewrite the want/ files of the snapshot archives in testdata") + +// templatesGo declares the templates variable for an archive that does not +// declare its own: every template file, parsed as ParseFS would. +const templatesGo = `package server + +import ( + "embed" + "html/template" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) +` + +// TestSnapshots generates the routes files for each archive in testdata +// and compares them with the archive's want/ files. +// +// An archive holds a case's inputs and what it generates, and snapshots, +// in snapshots_test.go, the configuration it is generated with: +// +// - Go files and .gohtml files are loaded as example.com/server by +// internal/load/loadtest -- type checked against the stub standard +// library, without the go command -- and hydrated by +// load.GenerateSource, as muxt generate does. An archive with no +// templates.go gets one declaring the templates variable over every +// .gohtml file. +// - want/ files are the expected output: one per generated file, named +// for it, and want/log.txt and want/error.txt for what generation +// logged and the error loading or generating returned. +// +// Paths in the output are relative to the directory the package was +// written to. +// +// Run with -update to rewrite the want/ files from the generator, then +// read the diff: the snapshot says what the generator does, not what it +// should do. Whether generated code compiles and serves requests is the +// integration suite's job, under cmd/muxt/testdata. +func TestSnapshots(t *testing.T) { + archives, err := filepath.Glob(filepath.Join("testdata", "*.txtar")) + if err != nil { + t.Fatal(err) + } + for _, archivePath := range archives { + name := strings.TrimSuffix(filepath.Base(archivePath), ".txtar") + if !slices.ContainsFunc(snapshots, func(c snapshotCase) bool { return c.archive == name }) { + t.Errorf("testdata/%s.txtar has no configuration in snapshots", name) + } + } + for _, tt := range snapshots { + t.Run(tt.archive, func(t *testing.T) { + archivePath := filepath.Join("testdata", tt.archive+".txtar") + archive, err := txtar.ParseFile(archivePath) + if err != nil { + t.Fatal(err) + } + got := snapshot(t, tt.config, archive) + if *update { + writeSnapshot(t, archivePath, archive, got) + return + } + want := make(map[string]string) + for _, file := range archive.Files { + if name, ok := strings.CutPrefix(file.Name, "want/"); ok { + want[name] = string(file.Data) + } + } + for _, name := range sortedKeys(got, want) { + if got[name] != want[name] { + t.Errorf("want/%s differs (run go test -run TestSnapshots -update to rewrite):\n--- got\n%s\n--- want\n%s", name, got[name], want[name]) + } + } + }) + } +} + +// snapshot loads an archive's package, generates from it, and returns the +// want/ files it produces, by name. +func snapshot(t *testing.T, config generate.RoutesFileConfiguration, archive *txtar.Archive) map[string]string { + t.Helper() + files := make(map[string]string) + for _, file := range archive.Files { + if strings.HasPrefix(file.Name, "want/") { + continue + } + files[file.Name] = string(file.Data) + } + if _, declared := files["templates.go"]; !declared { + files["templates.go"] = templatesGo + } + + dir := t.TempDir() + relative := func(text string) string { return strings.ReplaceAll(text, dir+string(filepath.Separator), "") } + got := make(map[string]string) + fail := func(err error) map[string]string { + text := err.Error() + if multiLine, ok := err.(muxt.MultiLineError); ok { + text = multiLine.MultiLineError() + } + got["error.txt"] = relative(text) + "\n" + return got + } + + pkg, receiver, err := load.GenerateSource(dir, loadtest.Package(t, dir, "example.com/server", files), config) + if err != nil { + return fail(err) + } + var logs strings.Builder + generated, err := generate.TemplateRoutesFiles(dir, config, pkg, receiver, log.New(&logs, "", 0)) + for _, file := range generated { + got[relative(file.Path)] = file.Content + for _, name := range unusedImports(t, file.Content) { + t.Errorf("%s imports %s without using it", relative(file.Path), name) + } + } + if logs.Len() > 0 { + got["log.txt"] = relative(logs.String()) + } + if err != nil { + return fail(err) + } + return got +} + +func writeSnapshot(t *testing.T, archivePath string, archive *txtar.Archive, got map[string]string) { + t.Helper() + files := slices.DeleteFunc(slices.Clone(archive.Files), func(file txtar.File) bool { + return strings.HasPrefix(file.Name, "want/") + }) + for _, name := range sortedKeys(got) { + files = append(files, txtar.File{Name: "want/" + name, Data: []byte(got[name])}) + } + archive.Files = files + if err := os.WriteFile(archivePath, txtar.Format(archive), 0o644); err != nil { + t.Fatal(err) + } +} + +func sortedKeys(maps ...map[string]string) []string { + var keys []string + for _, m := range maps { + for key := range m { + if !slices.Contains(keys, key) { + keys = append(keys, key) + } + } + } + slices.Sort(keys) + return keys +} + +// unusedImports names the imports a generated file declares but does not +// refer to. Generated code that compiles has none; a file that did would +// be a generator registering an import for a declaration it did not write. +func unusedImports(t *testing.T, content string) []string { + t.Helper() + file, err := parser.ParseFile(token.NewFileSet(), "generated.go", content, 0) + if err != nil { + t.Fatal(err) + } + referenced := make(map[string]bool) + ast.Inspect(file, func(node ast.Node) bool { + if sel, ok := node.(*ast.SelectorExpr); ok { + // A package name resolves to nothing in the file; a local + // variable of the same name does. + if id, ok := sel.X.(*ast.Ident); ok && id.Obj == nil { + referenced[id.Name] = true + } + } + return true + }) + var unused []string + for _, spec := range file.Imports { + importPath, err := strconv.Unquote(spec.Path.Value) + if err != nil { + t.Fatal(err) + } + name := path.Base(importPath) + if spec.Name != nil { + name = spec.Name.Name + } + if !referenced[name] { + unused = append(unused, importPath) + } + } + return unused +} diff --git a/internal/generate/snapshots_test.go b/internal/generate/snapshots_test.go new file mode 100644 index 00000000..2829470e --- /dev/null +++ b/internal/generate/snapshots_test.go @@ -0,0 +1,591 @@ +package generate_test + +import "github.com/typelate/muxt/internal/generate" + +// snapshots names the configuration each archive in testdata is generated +// with: the configuration the command line in the comment parses into. +// +// TestCommandLineConfigurations in internal/cli states what command lines +// parse into, with literals like these. Search for a literal to find its +// twin. They need not match field for field, but every configuration here +// is one a command line can produce: rejecting a configuration that cannot +// work is the command line's job, so nothing here tests one. +type snapshotCase struct { + archive string + config generate.RoutesFileConfiguration +} + +var snapshots = []snapshotCase{ + { + // muxt generate + archive: "err_duplicate_pattern", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate + archive: "err_name_errors", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "err_resolution", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "err_response_state_with_response_argument", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate + archive: "err_route_paths_method_collision", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate + archive: "err_signals_without_datastar", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "execute_callback", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=Server --output-file=routes.go --output-routes-func=Routes --output-receiver-interface=Handlers --output-template-data-type=Data --output-template-route-paths-type=Paths + archive: "flag_custom_names", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "Routes", + ReceiverType: "Server", + ReceiverInterface: "Handlers", + TemplateDataType: "Data", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "Paths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --output-htmx + archive: "flag_htmx", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputHTMX: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T --output-routes-func-with-logger-param --output-routes-func-with-path-prefix-param --output-routes-func-with-middleware-param + archive: "flag_logger_path_prefix_middleware", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + PathPrefix: true, + Logger: true, + Middleware: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T --output-multiple-files --output-routes-func-with-middleware-param + archive: "flag_multiple_files", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + Middleware: true, + OutputMultipleFiles: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --output-exported-default-identifiers=false + archive: "flag_unexported_identifiers", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "templateRoutes", + ReceiverInterface: "routesReceiver", + TemplateDataType: "templateData", + SSETemplateDataType: "sseTemplateData", + TemplateRoutePathsTypeName: "templateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --output-muxt-version=false + archive: "flag_without_muxt_version", + config: generate.RoutesFileConfiguration{ + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "form_struct", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "form_values", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate + archive: "inferred_methods", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "last_event_id", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "marshal_json", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T --output-multipart-max-memory=1MiB + archive: "multipart", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + MultipartMaxMemory: 1 << 20, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "nested_calls", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "path_parameter_types", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=Server + archive: "receiver_method_sets", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "Server", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "redirect", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "request_body", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "response_argument", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "result_shapes", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate + archive: "route_without_call", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T --output-datastar + archive: "sse_datastar", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputDatastar: true, + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "sse_messages_without_datastar", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "sse", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "status_codes", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, + { + // muxt generate --use-receiver-type=T + archive: "synthesized_method_note", + config: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "TemplateRoutes", + ReceiverType: "T", + ReceiverInterface: "RoutesReceiver", + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: "TemplateRoutePaths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputExportedDefaultIdentifiers: true, + OutputMuxtVersion: true, + }, + }, +} diff --git a/internal/generate/source_test.go b/internal/generate/source_test.go new file mode 100644 index 00000000..e319e9b8 --- /dev/null +++ b/internal/generate/source_test.go @@ -0,0 +1,56 @@ +package generate + +import ( + "go/types" + "html/template" + "testing" + + "github.com/typelate/muxt/internal/source" + "github.com/typelate/muxt/internal/typestest" +) + +// This file builds what generation reads without loading a module: a +// package type checked in memory against the stub standard library, +// and templates parsed from strings. + +// testConfig is the configuration muxt generate runs with when no flags +// are passed. +func testConfig() RoutesFileConfiguration { + return RoutesFileConfiguration{ + PackageName: "main", + OutputFileName: "template_routes.go", + RoutesFunction: DefaultRoutesFunctionName, + ReceiverInterface: DefaultReceiverInterfaceName, + TemplateDataType: "TemplateData", + SSETemplateDataType: "SSETemplateData", + TemplateRoutePathsTypeName: DefaultTemplateRoutePathsTypeName, + TemplatesVariables: []string{"templates"}, + OutputExportedDefaultIdentifiers: true, + } +} + +// testSource type checks goSource as example.com/server and parses +// templates into the templates variable. When receiverType is not empty +// it returns that type as the receiver, as --use-receiver-type would name +// it. +func testSource(t *testing.T, goSource, receiverType, templates string) (source.Package, *types.Named) { + t.Helper() + pkg := typestest.MustCheck(t, "example.com/server", goSource) + src := source.Package{ + Fset: typestest.FileSet, + Types: pkg, + Imports: typestest.Packages(), + Variables: []source.Variable{{ + Name: "templates", + Set: template.Must(template.New("templates").Parse(templates)), + }}, + } + if receiverType == "" { + return src, nil + } + obj := pkg.Scope().Lookup(receiverType) + if obj == nil { + t.Fatalf("source declares no %s", receiverType) + } + return src, obj.Type().(*types.Named) +} diff --git a/internal/generate/sse.go b/internal/generate/sse.go index 2cfa1c18..96a29fd9 100644 --- a/internal/generate/sse.go +++ b/internal/generate/sse.go @@ -81,7 +81,10 @@ func sseMethodHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.D validationFailureBlock := func(string) *ast.BlockStmt { return parseErrBlock() } // The result type is per-callback; arg parsing only needs ctx/lastEventID/path // (it ignores the result type), so pass an empty struct here. - body, err := appendParseArgumentStatements(body, def, file, types.NewStruct(nil, nil), sig, def.Arguments, nil, "", config, def.CallExpression(), validationFailureBlock, parseErrBlock) + // 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()) + body, err := appendParseArgumentStatements(body, def, file, types.NewStruct(nil, nil), sig, def.Arguments, nil, "", config, call, validationFailureBlock, parseErrBlock) if err != nil { return nil, err } @@ -114,7 +117,7 @@ func sseMethodHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.D }}}}, ) - callArgs := slices.Clone(def.CallExpression().Args) + callArgs := slices.Clone(call.Args) for i, arg := range def.Arguments { switch arg.Type { case muxt.ArgumentTypeExecute: diff --git a/internal/generate/template_route_path.go b/internal/generate/template_route_path.go index ef005a42..f02d8be7 100644 --- a/internal/generate/template_route_path.go +++ b/internal/generate/template_route_path.go @@ -59,7 +59,7 @@ func routePathTypeAndMethods(imports *File, config RoutesFileConfiguration, defs func routePathFunc(file *File, config RoutesFileConfiguration, def *muxt.Definition) (_ *ast.FuncDecl, usesEscaper, usesSegmentsEscaper bool, _ error) { const methodReceiverName = routePathsReceiverName - encodingPkg, ok := file.Types("encoding") + encodingPkg, ok := file.OutputPackage().Import("encoding") if !ok { return nil, false, false, fmt.Errorf(`the "encoding" package must be loaded`) } diff --git a/internal/generate/testdata/err_duplicate_pattern.txtar b/internal/generate/testdata/err_duplicate_pattern.txtar new file mode 100644 index 00000000..8100efa3 --- /dev/null +++ b/internal/generate/testdata/err_duplicate_pattern.txtar @@ -0,0 +1,10 @@ +Two templates may not register the same pattern. +-- index.gohtml -- +{{define "GET /a"}}{{end}} +{{define "GET /a"}}{{end}} +-- server.go -- +package server +-- want/error.txt -- +duplicate route pattern "GET /a" +index.gohtml:1:11: first defined here +index.gohtml:2:11: also defined here diff --git a/internal/generate/testdata/err_name_errors.txtar b/internal/generate/testdata/err_name_errors.txtar new file mode 100644 index 00000000..46192240 --- /dev/null +++ b/internal/generate/testdata/err_name_errors.txtar @@ -0,0 +1,19 @@ +Every malformed name is reported, each at its position. +-- index.gohtml -- +{{define "OPTIONS /a"}}{{end}} +{{define "GET /b//c"}}{{end}} +{{define "GET /{id}/{id} F(id)"}}{{end}} +-- server.go -- +package server +-- want/error.txt -- + OPTIONS /a + ^^^^^^^ +index.gohtml:1:11: OPTIONS method not allowed; allowed methods: GET, POST, PUT, PATCH, and DELETE + + GET /b//c + ^ +index.gohtml:2:18: path has an empty segment + + GET /{id}/{id} F(id) + ^^ +index.gohtml:3:22: path parameter name "id" is used more than once; parameter names must be unique within a path diff --git a/internal/generate/testdata/err_resolution.txtar b/internal/generate/testdata/err_resolution.txtar new file mode 100644 index 00000000..a17bdba1 --- /dev/null +++ b/internal/generate/testdata/err_resolution.txtar @@ -0,0 +1,31 @@ +Resolution errors from every route are reported together, pointing at +the argument or method at fault, with where the method is defined. +-- index.gohtml -- +{{define "GET /{id} Float(id)"}}{{end}} +{{define "GET /ctx Context(request)"}}{{end}} +{{define "GET /none NoResults()"}}{{end}} +-- server.go -- +package server + +import "context" + +type T struct{} + +func (T) Float(id float64) string { return "" } +func (T) Context(ctx context.Context) string { return "" } +func (T) NoResults() {} +-- want/error.txt -- + GET /ctx Context(request) + ^^^^^^^ +index.gohtml:2:28: method expects type context.Context but request is *http.Request +server.go:8:10: Context is defined here + + GET /none NoResults() + ^^^^^^^^^ +index.gohtml:3:21: method NoResults() has no results; it should have one or two +server.go:9:10: NoResults is defined here + + GET /{id} Float(id) + ^^ +index.gohtml:1:27: method param type float64 not supported (supported: string, bool, int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64, or a type whose pointer implements encoding.TextUnmarshaler; bind as string and parse it yourself for other values) +server.go:7:10: Float is defined here diff --git a/internal/generate/testdata/err_response_state_with_response_argument.txtar b/internal/generate/testdata/err_response_state_with_response_argument.txtar new file mode 100644 index 00000000..3c5b635d --- /dev/null +++ b/internal/generate/testdata/err_response_state_with_response_argument.txtar @@ -0,0 +1,14 @@ +A template may not set the status of a route whose method took the +response. +-- index.gohtml -- +{{define "GET /download Download(response)"}}{{.StatusCode 404}}{{end}} +-- server.go -- +package server + +import "net/http" + +type T struct{} + +func (T) Download(w http.ResponseWriter) string { return "" } +-- want/error.txt -- +index.gohtml:1:11: template "GET /download Download(response)" calls StatusCode but Download takes the http.ResponseWriter, so muxt writes no status code or redirect for this route: either drop the response argument or call response.WriteHeader in the method diff --git a/internal/generate/testdata/err_route_paths_method_collision.txtar b/internal/generate/testdata/err_route_paths_method_collision.txtar new file mode 100644 index 00000000..a39bc051 --- /dev/null +++ b/internal/generate/testdata/err_route_paths_method_collision.txtar @@ -0,0 +1,11 @@ +Handlers whose names export to the same TemplateRoutePaths method +collide. +-- index.gohtml -- +{{define "GET /a list()"}}{{end}} +{{define "GET /b List()"}}{{end}} +-- server.go -- +package server +-- want/error.txt -- +TemplateRoutePaths method name collision: handlers "list" and "List" both produce method "List" +index.gohtml:1:11: "list" is defined here +index.gohtml:2:11: "List" is defined here diff --git a/internal/generate/testdata/err_signals_without_datastar.txtar b/internal/generate/testdata/err_signals_without_datastar.txtar new file mode 100644 index 00000000..dff530af --- /dev/null +++ b/internal/generate/testdata/err_signals_without_datastar.txtar @@ -0,0 +1,7 @@ +signals is Datastar's request body, so it needs --output-datastar. +-- index.gohtml -- +{{define "POST /count Count(signals)"}}{{end}} +-- server.go -- +package server +-- want/error.txt -- +the signals argument in "POST /count Count(signals)" requires --output-datastar; it is shorthand for unmarshalJSON(body) diff --git a/internal/generate/testdata/execute_callback.txtar b/internal/generate/testdata/execute_callback.txtar new file mode 100644 index 00000000..e182c453 --- /dev/null +++ b/internal/generate/testdata/execute_callback.txtar @@ -0,0 +1,199 @@ +A method taking execute renders the template when it calls back, with +the data it passes or none. +-- index.gohtml -- +{{define "GET /with-data Render(ctx, execute)"}}{{.Result}}{{end}} +{{define "GET /without-data Plain(execute)"}}{{end}} +-- server.go -- +package server + +import "context" + +type T struct{} + +func (T) Render(ctx context.Context, execute func(string) error) error { return nil } +func (T) Plain(execute func() error) error { return nil } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "context" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" + "sync/atomic" +) + +type RoutesReceiver interface { + Render(ctx context.Context, execute func(string) error) error + Plain(execute func() error) error +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /with-data", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + ctx := request.Context() + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + var executed atomic.Bool + if len(td.errList) == 0 { + if err := receiver.Render(ctx, func(data string) error { + if !executed.CompareAndSwap(false, true) { + return errors.New("execute callback called more than once") + } + td.result = data + return templates.ExecuteTemplate(buf, "GET /with-data Render(ctx, execute)", &td) + }); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + td.okay = true + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /without-data", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, struct{}]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + var executed atomic.Bool + if len(td.errList) == 0 { + if err := receiver.Plain(func() error { + if !executed.CompareAndSwap(false, true) { + return errors.New("execute callback called more than once") + } + return templates.ExecuteTemplate(buf, "GET /without-data Plain(execute)", &td) + }); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + td.okay = true + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Render() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "with-data") +} + +func (routePaths TemplateRoutePaths) Plain() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "without-data") +} diff --git a/internal/generate/testdata/flag_custom_names.txtar b/internal/generate/testdata/flag_custom_names.txtar new file mode 100644 index 00000000..b293c848 --- /dev/null +++ b/internal/generate/testdata/flag_custom_names.txtar @@ -0,0 +1,161 @@ +The output flags rename the generated file and identifiers. +-- index.gohtml -- +{{define "GET /{id} Show(id)"}}{{.Result}}{{end}} +-- server.go -- +package server + +type Server struct{} + +func (*Server) Show(id string) string { return id } +-- want/routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "net/url" + "path" + "strconv" + "sync" +) + +type Handlers interface { + Show(id string) string +} + +func Routes(mux *http.ServeMux, receiver Handlers) Paths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /{id}", func(response http.ResponseWriter, request *http.Request) { + var td = Data[Handlers, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + idPathParam := request.PathValue("id") + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Show(idPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /{id} Show(id)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return Paths{pathsPrefix: pathsPrefix} +} + +type Data[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *Data[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *Data[R, T]) Path() Paths { + return Paths{pathsPrefix: data.pathsPrefix} +} + +func (data *Data[R, T]) Result() T { + return data.result +} + +func (data *Data[R, T]) Request() *http.Request { + return data.request +} + +func (data *Data[R, T]) StatusCode(statusCode int) *Data[R, T] { + data.statusCode = statusCode + return data +} + +func (data *Data[R, T]) Header(key, value string) *Data[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *Data[R, T]) Ok() bool { + return data.okay +} + +func (data *Data[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *Data[R, T]) Receiver() R { + return data.receiver +} + +func (data *Data[R, T]) Redirect(url string, code int) (*Data[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *Data[R, T]) RedirectMultipleChoices(url string) (*Data[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *Data[R, T]) RedirectMovedPermanently(url string) (*Data[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *Data[R, T]) RedirectFound(url string) (*Data[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *Data[R, T]) RedirectSeeOther(url string) (*Data[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *Data[R, T]) String() string { + return "" +} + +type Paths struct { + pathsPrefix string +} + +func (routePaths Paths) Show(idPathParam string) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), routePaths.escapePathSegment(idPathParam)) +} + +func (routePaths Paths) escapePathSegment(value string) string { + switch value { + case ".": + return "%2E" + case "..": + return "%2E%2E" + } + return url.PathEscape(value) +} diff --git a/internal/generate/testdata/flag_htmx.txtar b/internal/generate/testdata/flag_htmx.txtar new file mode 100644 index 00000000..0b83b313 --- /dev/null +++ b/internal/generate/testdata/flag_htmx.txtar @@ -0,0 +1,216 @@ +--output-htmx adds the HX header helpers to TemplateData. +-- index.gohtml -- +{{define "GET /{$}"}}home{{end}} +-- server.go -- +package server +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "strconv" + "sync" +) + +type RoutesReceiver interface { +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /{$}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, struct { + }]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if err := templates.ExecuteTemplate(buf, "GET /{$}", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +func (data *TemplateData[R, T]) HXLocation(link string) *TemplateData[R, T] { + return data.Header("HX-Location", link) +} + +func (data *TemplateData[R, T]) HXPushURL(link string) *TemplateData[R, T] { + return data.Header("HX-Push-Url", link) +} + +func (data *TemplateData[R, T]) HXRedirect(link string) *TemplateData[R, T] { + return data.Header("HX-Redirect", link) +} + +func (data *TemplateData[R, T]) HXRefresh() *TemplateData[R, T] { + return data.Header("HX-Refresh", "true") +} + +func (data *TemplateData[R, T]) HXReplaceURL(link string) *TemplateData[R, T] { + return data.Header("HX-Replace-Url", link) +} + +func (data *TemplateData[R, T]) HXReswap(swap string) *TemplateData[R, T] { + return data.Header("HX-Reswap", swap) +} + +func (data *TemplateData[R, T]) HXRetarget(target string) *TemplateData[R, T] { + return data.Header("HX-Retarget", target) +} + +func (data *TemplateData[R, T]) HXReselect(selector string) *TemplateData[R, T] { + return data.Header("HX-Reselect", selector) +} + +func (data *TemplateData[R, T]) HXTrigger(eventName string) *TemplateData[R, T] { + return data.Header("HX-Trigger", eventName) +} + +func (data *TemplateData[R, T]) HXTriggerAfterSettle(eventName string) *TemplateData[R, T] { + return data.Header("HX-Trigger-After-Settle", eventName) +} + +func (data *TemplateData[R, T]) HXTriggerAfterSwap(eventName string) *TemplateData[R, T] { + return data.Header("HX-Trigger-After-Swap", eventName) +} + +func (data *TemplateData[R, T]) HXBoosted() bool { + return data.Request().Header.Get("HX-Boosted") != "" +} + +func (data *TemplateData[R, T]) HXCurrentURL() string { + return data.Request().Header.Get("HX-Current-Url") +} + +func (data *TemplateData[R, T]) HXHistoryRestoreRequest() bool { + return data.Request().Header.Get("HX-History-Restore-Request") == "true" +} + +func (data *TemplateData[R, T]) HXPrompt() string { + return data.Request().Header.Get("HX-Prompt") +} + +func (data *TemplateData[R, T]) HXRequest() bool { + return data.Request().Header.Get("HX-Request") == "true" +} + +func (data *TemplateData[R, T]) HXTargetElementID() string { + return data.Request().Header.Get("HX-Target") +} + +func (data *TemplateData[R, T]) HXTriggerName() string { + return data.Request().Header.Get("HX-Trigger-Name") +} + +func (data *TemplateData[R, T]) HXTriggerElementID() string { + return data.Request().Header.Get("HX-Trigger") +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) ReadExact() string { + return "/" +} diff --git a/internal/generate/testdata/flag_logger_path_prefix_middleware.txtar b/internal/generate/testdata/flag_logger_path_prefix_middleware.txtar new file mode 100644 index 00000000..94c17289 --- /dev/null +++ b/internal/generate/testdata/flag_logger_path_prefix_middleware.txtar @@ -0,0 +1,189 @@ +The routes function takes a logger, a path prefix, and a middleware. +-- index.gohtml -- +{{define "GET /{$}"}}home{{end}} +{{define "GET /article/{id} Article(id)"}}{{.Result}}{{end}} +-- server.go -- +package server + +type T struct{} + +func (T) Article(id int) string { return "" } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Article(id int) string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver, logger *slog.Logger, pathsPrefix string, middleware func(next http.Handler) http.Handler) TemplateRoutePaths { + if middleware == nil { + middleware = func(next http.Handler) http.Handler { + return next + } + } + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.Handle("GET "+path.Join(pathsPrefix, "/article/{id}"), middleware(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + idParsed, err := strconv.Atoi(request.PathValue("id")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + idPathParam := idParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Article(idPathParam) + td.okay = true + } + logger.DebugContext(request.Context(), "handling request", slog.String("pattern", "GET /article/{id}"), slog.String("path", request.URL.Path), slog.String("method", request.Method)) + if err := templates.ExecuteTemplate(buf, "GET /article/{id} Article(id)", &td); err != nil { + logger.ErrorContext(request.Context(), "failed to render page", slog.String("pattern", "GET /article/{id}"), slog.String("path", request.URL.Path), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }))) + mux.Handle("GET "+path.Join(pathsPrefix, "/{$}"), middleware(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, struct { + }]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + logger.DebugContext(request.Context(), "handling request", slog.String("pattern", "GET /{$}"), slog.String("path", request.URL.Path), slog.String("method", request.Method)) + if err := templates.ExecuteTemplate(buf, "GET /{$}", &td); err != nil { + logger.ErrorContext(request.Context(), "failed to render page", slog.String("pattern", "GET /{$}"), slog.String("path", request.URL.Path), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }))) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Article(idPathParam int) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "article", strconv.Itoa(idPathParam)) +} + +func (routePaths TemplateRoutePaths) ReadExact() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/")) +} diff --git a/internal/generate/testdata/flag_multiple_files.txtar b/internal/generate/testdata/flag_multiple_files.txtar new file mode 100644 index 00000000..b5d7053c --- /dev/null +++ b/internal/generate/testdata/flag_multiple_files.txtar @@ -0,0 +1,280 @@ +--output-multiple-files writes the routes for each template file into +its own file; routes from templates with no file stay in the main one. +-- index.gohtml -- +{{define "GET /{$} Home()"}}{{.Result}}{{end}} +-- user-profile.gohtml -- +{{define "GET /user/{id} User(id)"}}{{.Result}}{{end}} +{{define "POST /user/{id} Save(id, form)"}}{{.Result}}{{end}} +-- server.go -- +package server + +import "net/url" + +type T struct{} + +func (T) Home() string { return "" } +func (T) User(id int) string { return "" } +func (T) Save(id int, form url.Values) error { return nil } +-- want/index_template_routes_gen.go -- +package server + +import ( + "bytes" + "cmp" + "log/slog" + "net/http" + "strconv" + "sync" +) + +type indexRoutesReceiver interface { + Home() string +} + +func indexTemplateRoutes(mux *http.ServeMux, receiver indexRoutesReceiver, pathsPrefix string, middleware func(next http.Handler) http.Handler) { + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.Handle("GET /{$}", middleware(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[indexRoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Home() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /{$} Home()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }))) +} +-- want/template_routes.go -- +package server + +import ( + "cmp" + "errors" + "fmt" + "net/http" + "path" + "strconv" +) + +type RoutesReceiver interface { + indexRoutesReceiver + userProfileRoutesReceiver +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver, middleware func(next http.Handler) http.Handler) TemplateRoutePaths { + pathsPrefix := "" + if middleware == nil { + middleware = func(next http.Handler) http.Handler { + return next + } + } + indexTemplateRoutes(mux, receiver, pathsPrefix, middleware) + userProfileTemplateRoutes(mux, receiver, pathsPrefix, middleware) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) User(idPathParam int) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "user", strconv.Itoa(idPathParam)) +} + +func (routePaths TemplateRoutePaths) Save(idPathParam int) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "user", strconv.Itoa(idPathParam)) +} + +func (routePaths TemplateRoutePaths) Home() string { + return "/" +} +-- want/user-profile_template_routes_gen.go -- +package server + +import ( + "bytes" + "cmp" + "log/slog" + "net/http" + "net/url" + "strconv" + "sync" +) + +type userProfileRoutesReceiver interface { + User(id int) string + Save(id int, form url.Values) error +} + +func userProfileTemplateRoutes(mux *http.ServeMux, receiver userProfileRoutesReceiver, pathsPrefix string, middleware func(next http.Handler) http.Handler) { + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.Handle("GET /user/{id}", middleware(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[userProfileRoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + idParsed, err := strconv.Atoi(request.PathValue("id")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + idPathParam := idParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.User(idPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /user/{id} User(id)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }))) + mux.Handle("POST /user/{id}", middleware(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[userProfileRoutesReceiver, error]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + idParsed, err := strconv.Atoi(request.PathValue("id")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + idPathParam := idParsed + if err := request.ParseForm(); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + var form url.Values = request.Form + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Save(idPathParam, form) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /user/{id} Save(id, form)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }))) +} diff --git a/internal/generate/testdata/flag_unexported_identifiers.txtar b/internal/generate/testdata/flag_unexported_identifiers.txtar new file mode 100644 index 00000000..f8eee837 --- /dev/null +++ b/internal/generate/testdata/flag_unexported_identifiers.txtar @@ -0,0 +1,11 @@ +--output-exported-default-identifiers=false names the generated +identifiers unexported, as the command line spells them. +-- index.gohtml -- +{{define "GET /events sse(Events(execute))"}}{{.Result}}{{end}} +{{define "GET /{$}"}}home{{end}} +-- server.go -- +package server +-- want/error.txt -- + GET /events sse(Events(execute)) + ^^^^^^^ +index.gohtml:1:34: method Events using the execute callback must be defined on the receiver type diff --git a/internal/generate/testdata/flag_without_muxt_version.txtar b/internal/generate/testdata/flag_without_muxt_version.txtar new file mode 100644 index 00000000..b326c0d7 --- /dev/null +++ b/internal/generate/testdata/flag_without_muxt_version.txtar @@ -0,0 +1,136 @@ +With --output-muxt-version=false, TemplateData has no MuxtVersion method, +so a template calling it does not compile. +-- index.gohtml -- +{{define "GET /{$}"}}{{.MuxtVersion}}{{end}} +-- server.go -- +package server +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "strconv" + "sync" +) + +type RoutesReceiver interface { +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /{$}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, struct { + }]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if err := templates.ExecuteTemplate(buf, "GET /{$}", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) ReadExact() string { + return "/" +} diff --git a/internal/generate/testdata/form_struct.txtar b/internal/generate/testdata/form_struct.txtar new file mode 100644 index 00000000..02c10cf0 --- /dev/null +++ b/internal/generate/testdata/form_struct.txtar @@ -0,0 +1,230 @@ +A form struct binds each field from the form, parsing scalars and +slices, renaming inputs with the name tag, and validating from the +input a template tag names. +-- index.gohtml -- +{{define "POST /signup Signup(form)"}}{{end}} +{{define "signup-form"}} + + +{{end}} +-- server.go -- +package server + +type Signup struct { + Age int `name:"age" template:"signup-form"` + Handle string `name:"handle" template:"signup-form"` + Tags []string `name:"tag"` + Score float64 + Verified bool + Count uint32 +} + +type T struct{} + +func (T) Signup(form Signup) error { return nil } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "regexp" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Signup(form Signup) error +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("POST /signup", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, error]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + if err := request.ParseForm(); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + var form Signup + { + value, err := strconv.Atoi(request.FormValue("age")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } else { + if value < 13 { + td.errList = append(td.errList, errors.New("age must not be less than 13")) + td.errStatusCode = http.StatusBadRequest + } + if value > 130 { + td.errList = append(td.errList, errors.New("age must not be more than 130")) + td.errStatusCode = http.StatusBadRequest + } + } + form.Age = value + } + { + value := request.FormValue("handle") + if !regexp.MustCompile("[a-z]+").MatchString(value) { + td.errList = append(td.errList, errors.New("handle must match \"[a-z]+\"")) + td.errStatusCode = http.StatusBadRequest + } + if len(value) < 2 { + td.errList = append(td.errList, errors.New("handle is too short (the min length is 2)")) + td.errStatusCode = http.StatusBadRequest + } + if len(value) > 20 { + td.errList = append(td.errList, errors.New("handle is too long (the max length is 20)")) + td.errStatusCode = http.StatusBadRequest + } + form.Handle = value + } + for _, val := range request.Form["tag"] { + form.Tags = append(form.Tags, val) + } + { + value, err := strconv.ParseFloat(request.FormValue("Score"), 64) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + form.Score = value + } + { + value, err := strconv.ParseBool(request.FormValue("Verified")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + form.Verified = value + } + { + value, err := strconv.ParseUint(request.FormValue("Count"), 10, 32) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + form.Count = uint32(value) + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Signup(form) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /signup Signup(form)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Signup() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "signup") +} diff --git a/internal/generate/testdata/form_values.txtar b/internal/generate/testdata/form_values.txtar new file mode 100644 index 00000000..450bbe69 --- /dev/null +++ b/internal/generate/testdata/form_values.txtar @@ -0,0 +1,193 @@ +A url.Values parameter receives the parsed form as it is. +-- index.gohtml -- +{{define "POST /values Values(form)"}}{{end}} +{{define "POST /unmarshal Values(unmarshalForm(body))"}}{{end}} +-- server.go -- +package server + +import "net/url" + +type T struct{} + +func (T) Values(form url.Values) string { return "" } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "net/url" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Values(form url.Values) string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("POST /unmarshal", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + if err := request.ParseForm(); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + var form url.Values = request.Form + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Values(form) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /unmarshal Values(unmarshalForm(body))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("POST /values", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + if err := request.ParseForm(); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + var form url.Values = request.Form + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Values(form) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /values Values(form)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) CreateUnmarshalCallingValues() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "unmarshal") +} + +func (routePaths TemplateRoutePaths) CreateValuesCallingValues() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "values") +} diff --git a/internal/generate/testdata/inferred_methods.txtar b/internal/generate/testdata/inferred_methods.txtar new file mode 100644 index 00000000..f7c60ade --- /dev/null +++ b/internal/generate/testdata/inferred_methods.txtar @@ -0,0 +1,222 @@ +Without --use-receiver-type every called method is inferred from its +arguments, and returns any. +-- index.gohtml -- +{{define "GET /{id} Get(ctx, id)"}}{{.Result}}{{end}} +{{define "POST /upload Upload(request, response)"}}{{end}} +{{define "PATCH /note Note(form, lastEventID)"}}{{end}} +-- server.go -- +package server +-- want/log.txt -- +warning: POST /upload uses the response argument, so muxt does not manage this route's status codes, headers, or rendering; silence with MUXT_SILENCE_WARNING_HTTP_RESPONSE_ARGUMENT=true +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "context" + "errors" + "fmt" + "log/slog" + "net/http" + "net/url" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Note(form url.Values, lastEventID string) any + Upload(request *http.Request, response http.ResponseWriter) any + Get(ctx context.Context, id string) any +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("PATCH /note", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, any]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + if err := request.ParseForm(); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + var form url.Values = request.Form + lastEventID := request.Header.Get("Last-Event-Id") + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Note(form, lastEventID) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "PATCH /note Note(form, lastEventID)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("POST /upload", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, any]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Upload(request, response) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /upload Upload(request, response)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /{id}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, any]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + ctx := request.Context() + idPathParam := request.PathValue("id") + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Get(ctx, idPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /{id} Get(ctx, id)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Note() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "note") +} + +func (routePaths TemplateRoutePaths) Upload() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "upload") +} + +func (routePaths TemplateRoutePaths) Get(idPathParam string) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), routePaths.escapePathSegment(idPathParam)) +} + +func (routePaths TemplateRoutePaths) escapePathSegment(value string) string { + switch value { + case ".": + return "%2E" + case "..": + return "%2E%2E" + } + return url.PathEscape(value) +} diff --git a/internal/generate/testdata/last_event_id.txtar b/internal/generate/testdata/last_event_id.txtar new file mode 100644 index 00000000..2394a7b5 --- /dev/null +++ b/internal/generate/testdata/last_event_id.txtar @@ -0,0 +1,156 @@ +lastEventID reads the Last-Event-Id header, parsed into the parameter +type. +-- index.gohtml -- +{{define "GET /resume Resume(lastEventID)"}}{{.Result}}{{end}} +-- server.go -- +package server + +type T struct{} + +func (T) Resume(lastEventID uint64) string { return "" } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Resume(lastEventID uint64) string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /resume", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + lastEventIDParsed, err := strconv.ParseUint(request.Header.Get("Last-Event-Id"), 10, 64) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + lastEventID := lastEventIDParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Resume(lastEventID) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /resume Resume(lastEventID)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Resume() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "resume") +} diff --git a/internal/generate/testdata/marshal_json.txtar b/internal/generate/testdata/marshal_json.txtar new file mode 100644 index 00000000..4dea7229 --- /dev/null +++ b/internal/generate/testdata/marshal_json.txtar @@ -0,0 +1,173 @@ +marshalJSON writes the result as JSON, rendering the template for its +side effects and on errors. +-- index.gohtml -- +{{define "GET /api/{id} marshalJSON(Get(id))"}}{{.Err}}{{end}} +-- server.go -- +package server + +type Item struct{ ID int } + +type T struct{} + +func (T) Get(id int) (Item, error) { return Item{}, nil } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "encoding/json" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Get(id int) (Item, error) +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /api/{id}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, Item]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + idParsed, err := strconv.Atoi(request.PathValue("id")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + idPathParam := idParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + var err error + td.result, err = receiver.Get(idPathParam) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusInternalServerError + } + } + if err := templates.ExecuteTemplate(buf, "GET /api/{id} marshalJSON(Get(id))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + if len(td.errList) == 0 { + buf.Reset() + jsonBody, err := json.Marshal(td.result) + if err != nil { + http.Error(response, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) + return + } + _, _ = buf.Write(jsonBody) + response.Header().Set("content-type", "application/json") + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Get(idPathParam int) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "api", strconv.Itoa(idPathParam)) +} diff --git a/internal/generate/testdata/multipart.txtar b/internal/generate/testdata/multipart.txtar new file mode 100644 index 00000000..62eaf82c --- /dev/null +++ b/internal/generate/testdata/multipart.txtar @@ -0,0 +1,211 @@ +A multipart struct binds text fields and file headers; a +*multipart.Form parameter receives the parsed form. +-- index.gohtml -- +{{define "POST /upload Upload(multipart)"}}{{end}} +{{define "POST /raw Raw(multipart)"}}{{end}} +-- server.go -- +package server + +import "mime/multipart" + +type Upload struct { + Title string + File *multipart.FileHeader + Files []*multipart.FileHeader `name:"attachment"` +} + +type T struct{} + +func (T) Upload(in Upload) error { return nil } +func (T) Raw(form *multipart.Form) error { return nil } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "mime/multipart" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Raw(form *multipart.Form) error + Upload(in Upload) error +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("POST /raw", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, error]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + if err := request.ParseMultipartForm(1048576); err != nil && !errors.Is(err, http.ErrNotMultipart) { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + var multipart *multipart.Form = request.MultipartForm + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Raw(multipart) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /raw Raw(multipart)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("POST /upload", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, error]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + if err := request.ParseMultipartForm(1048576); err != nil && !errors.Is(err, http.ErrNotMultipart) { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + var multipart Upload + multipart.Title = request.FormValue("Title") + if request.MultipartForm != nil { + if fhs := request.MultipartForm.File["File"]; len(fhs) > 0 { + multipart.File = fhs[0] + } + } + if request.MultipartForm != nil { + multipart.Files = request.MultipartForm.File["attachment"] + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Upload(multipart) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /upload Upload(multipart)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Raw() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "raw") +} + +func (routePaths TemplateRoutePaths) Upload() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "upload") +} diff --git a/internal/generate/testdata/nested_calls.txtar b/internal/generate/testdata/nested_calls.txtar new file mode 100644 index 00000000..a96852b1 --- /dev/null +++ b/internal/generate/testdata/nested_calls.txtar @@ -0,0 +1,210 @@ +A call's argument may be the result of another call, a receiver method +or a package function, which runs first. +-- index.gohtml -- +{{define "GET /article/{id} Show(ctx, Load(ctx, id))"}}{{.Result}}{{end}} +{{define "GET /user Show(ctx, Current(request))"}}{{.Result}}{{end}} +-- server.go -- +package server + +import ( + "context" + "net/http" +) + +type Article struct{ Title string } + +type T struct{} + +func (T) Load(ctx context.Context, id int) (Article, error) { return Article{}, nil } +func (T) Show(ctx context.Context, a Article) string { return "" } + +func Current(r *http.Request) (Article, bool) { return Article{}, false } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "context" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Load(ctx context.Context, id int) (Article, error) + Show(ctx context.Context, a Article) string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /article/{id}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + ctx := request.Context() + idParsed, err := strconv.Atoi(request.PathValue("id")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + idPathParam := idParsed + result0, err := receiver.Load(ctx, idPathParam) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusInternalServerError + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Show(ctx, result0) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /article/{id} Show(ctx, Load(ctx, id))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /user", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + ctx := request.Context() + result0, ok := Current(request) + if !ok { + return + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Show(ctx, result0) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /user Show(ctx, Current(request))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) ReadArticleByIDCallingShow(idPathParam int) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "article", strconv.Itoa(idPathParam)) +} + +func (routePaths TemplateRoutePaths) ReadUserCallingShow() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "user") +} diff --git a/internal/generate/testdata/path_parameter_types.txtar b/internal/generate/testdata/path_parameter_types.txtar new file mode 100644 index 00000000..fd13d1cb --- /dev/null +++ b/internal/generate/testdata/path_parameter_types.txtar @@ -0,0 +1,413 @@ +Each path parameter parses into the type of the parameter it is passed +to, and the route path helper takes that type. +-- index.gohtml -- +{{define "GET /int/{id} Int(id)"}}{{.Result}}{{end}} +{{define "GET /small/{a}/{b}/{c} Small(a, b, c)"}}{{.Result}}{{end}} +{{define "GET /bool/{flag} Bool(flag)"}}{{.Result}}{{end}} +{{define "GET /string/{name} String(name)"}}{{.Result}}{{end}} +{{define "GET /at/{at} At(at)"}}{{.Result}}{{end}} +{{define "GET /files/{path...} Files(path)"}}{{.Result}}{{end}} +{{define "GET /exact/{$} Exact()"}}{{.Result}}{{end}} +-- server.go -- +package server + +import "time" + +type T struct{} + +func (T) Int(id int) int { return id } +func (T) Small(a int8, b uint16, c int64) int64 { return c } +func (T) Bool(flag bool) bool { return flag } +func (T) String(name string) string { return name } +func (T) At(at time.Time) time.Time { return at } +func (T) Files(path string) string { return path } +func (T) Exact() string { return "" } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "net/url" + "path" + "strconv" + "strings" + "sync" + "time" +) + +type RoutesReceiver interface { + At(at time.Time) time.Time + Bool(flag bool) bool + Exact() string + Files(path string) string + Int(id int) int + Small(a int8, b uint16, c int64) int64 + String(name string) string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /at/{at}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, time.Time]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + var atParsed time.Time + if err := atParsed.UnmarshalText([]byte(request.PathValue("at"))); err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + atPathParam := atParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.At(atPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /at/{at} At(at)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /bool/{flag}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, bool]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + flagParsed, err := strconv.ParseBool(request.PathValue("flag")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + flagPathParam := flagParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Bool(flagPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /bool/{flag} Bool(flag)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /exact/{$}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Exact() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /exact/{$} Exact()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /files/{path...}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + pathPathParam := request.PathValue("path") + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Files(pathPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /files/{path...} Files(path)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /int/{id}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, int]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + idParsed, err := strconv.Atoi(request.PathValue("id")) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + idPathParam := idParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Int(idPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /int/{id} Int(id)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /small/{a}/{b}/{c}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, int64]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + aParsed, err := strconv.ParseInt(request.PathValue("a"), 10, 8) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + aPathParam := int8(aParsed) + bParsed, err := strconv.ParseUint(request.PathValue("b"), 10, 16) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + bPathParam := uint16(bParsed) + cParsed, err := strconv.ParseInt(request.PathValue("c"), 10, 64) + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + cPathParam := cParsed + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Small(aPathParam, bPathParam, cPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /small/{a}/{b}/{c} Small(a, b, c)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /string/{name}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + namePathParam := request.PathValue("name") + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.String(namePathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /string/{name} String(name)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) At(atPathParam time.Time) (string, error) { + segment2_06e81700, err := atPathParam.MarshalText() + if err != nil { + return "", fmt.Errorf("failed to marshal path value {at} (segment 2) in /at/{at}: %w", err) + } + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "at", routePaths.escapePathSegment(string(segment2_06e81700))), nil +} + +func (routePaths TemplateRoutePaths) Bool(flagPathParam bool) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "bool", strconv.FormatBool(bool(flagPathParam))) +} + +func (routePaths TemplateRoutePaths) Exact() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "exact") + "/" +} + +func (routePaths TemplateRoutePaths) Files(pathPathParam string) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "files", routePaths.escapePathSegments(pathPathParam)) +} + +func (routePaths TemplateRoutePaths) Int(idPathParam int) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "int", strconv.Itoa(idPathParam)) +} + +func (routePaths TemplateRoutePaths) Small(aPathParam int8, bPathParam uint16, cPathParam int64) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "small", strconv.FormatInt(int64(aPathParam), 10), strconv.FormatUint(uint64(bPathParam), 10), strconv.FormatInt(int64(cPathParam), 10)) +} + +func (routePaths TemplateRoutePaths) String(namePathParam string) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "string", routePaths.escapePathSegment(namePathParam)) +} + +func (routePaths TemplateRoutePaths) escapePathSegment(value string) string { + switch value { + case ".": + return "%2E" + case "..": + return "%2E%2E" + } + return url.PathEscape(value) +} + +func (routePaths TemplateRoutePaths) escapePathSegments(value string) string { + segments := strings.Split(value, "/") + for i, segment := range segments { + segments[i] = routePaths.escapePathSegment(segment) + } + return strings.Join(segments, "/") +} diff --git a/internal/generate/testdata/receiver_method_sets.txtar b/internal/generate/testdata/receiver_method_sets.txtar new file mode 100644 index 00000000..c0770df4 --- /dev/null +++ b/internal/generate/testdata/receiver_method_sets.txtar @@ -0,0 +1,220 @@ +Methods are found through embedded fields and pointer receivers; +functions in the package are called directly and never join the +receiver interface. +-- index.gohtml -- +{{define "GET /embedded Embedded()"}}{{.Result}}{{end}} +{{define "GET /pointer Pointer()"}}{{.Result}}{{end}} +{{define "GET /function Function()"}}{{.Result}}{{end}} +-- server.go -- +package server + +type Base struct{} + +func (Base) Embedded() string { return "" } + +type Server struct{ Base } + +func (*Server) Pointer() int { return 0 } + +func Function() bool { return false } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Embedded() string + Pointer() int +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /embedded", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Embedded() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /embedded Embedded()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /function", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, bool]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = Function() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /function Function()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /pointer", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, int]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Pointer() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /pointer Pointer()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Embedded() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "embedded") +} + +func (routePaths TemplateRoutePaths) Function() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "function") +} + +func (routePaths TemplateRoutePaths) Pointer() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "pointer") +} diff --git a/internal/generate/testdata/redirect.txtar b/internal/generate/testdata/redirect.txtar new file mode 100644 index 00000000..b825cadc --- /dev/null +++ b/internal/generate/testdata/redirect.txtar @@ -0,0 +1,220 @@ +Only a template that may call a redirect method gets the redirect block. +-- index.gohtml -- +{{define "POST /redirects Save()"}}{{.RedirectSeeOther "/"}}{{end}} +{{define "POST /via-partial Save()"}}{{template "partial" .}}{{end}} +{{define "partial"}}{{.Redirect "/" 302}}{{end}} +{{define "GET /stays Save()"}}{{.Result}}{{end}} +-- server.go -- +package server + +type T struct{} + +func (T) Save() string { return "" } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Save() string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("POST /redirects", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Save() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /redirects Save()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if td.redirectURL != "" { + http.Redirect(response, request, td.redirectURL, statusCode) + return + } + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /stays", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Save() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /stays Save()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("POST /via-partial", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Save() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /via-partial Save()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if td.redirectURL != "" { + http.Redirect(response, request, td.redirectURL, statusCode) + return + } + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) CreateRedirectsCallingSave() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "redirects") +} + +func (routePaths TemplateRoutePaths) ReadStaysCallingSave() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "stays") +} + +func (routePaths TemplateRoutePaths) CreateViaPartialCallingSave() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "via-partial") +} diff --git a/internal/generate/testdata/request_body.txtar b/internal/generate/testdata/request_body.txtar new file mode 100644 index 00000000..1467b406 --- /dev/null +++ b/internal/generate/testdata/request_body.txtar @@ -0,0 +1,200 @@ +The body is passed as an io.Reader, or decoded as JSON into the +parameter's type. +-- index.gohtml -- +{{define "POST /raw Raw(body)"}}{{end}} +{{define "POST /json Decode(ctx, unmarshalJSON(body))"}}{{end}} +-- server.go -- +package server + +import ( + "context" + "io" +) + +type Payload struct{ Name string } + +type T struct{} + +func (T) Raw(body io.Reader) error { return nil } +func (T) Decode(ctx context.Context, in Payload) error { return nil } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Decode(ctx context.Context, in Payload) error + Raw(body io.Reader) error +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("POST /json", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, error]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + ctx := request.Context() + var bodyValue Payload + if err := json.NewDecoder(request.Body).Decode(&bodyValue); err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusBadRequest + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Decode(ctx, bodyValue) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /json Decode(ctx, unmarshalJSON(body))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("POST /raw", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, error]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + body := request.Body + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Raw(body) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /raw Raw(body)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Decode() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "json") +} + +func (routePaths TemplateRoutePaths) Raw() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "raw") +} diff --git a/internal/generate/testdata/response_argument.txtar b/internal/generate/testdata/response_argument.txtar new file mode 100644 index 00000000..c6032027 --- /dev/null +++ b/internal/generate/testdata/response_argument.txtar @@ -0,0 +1,142 @@ +A method taking the response writes it itself: muxt writes no status. +-- index.gohtml -- +{{define "GET /download Download(response, request)"}}{{end}} +-- server.go -- +package server + +import "net/http" + +type T struct{} + +func (T) Download(w http.ResponseWriter, r *http.Request) string { return "" } +-- want/log.txt -- +warning: GET /download uses the response argument, so muxt does not manage this route's status codes, headers, or rendering; silence with MUXT_SILENCE_WARNING_HTTP_RESPONSE_ARGUMENT=true +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "sync" +) + +type RoutesReceiver interface { + Download(w http.ResponseWriter, r *http.Request) string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /download", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Download(response, request) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /download Download(response, request)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Download() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "download") +} diff --git a/internal/generate/testdata/result_shapes.txtar b/internal/generate/testdata/result_shapes.txtar new file mode 100644 index 00000000..b43194b4 --- /dev/null +++ b/internal/generate/testdata/result_shapes.txtar @@ -0,0 +1,256 @@ +A method returns a value, a value and an error, or a value and a bool. +-- index.gohtml -- +{{define "GET /value Value()"}}{{.Result}}{{end}} +{{define "GET /error ValueError()"}}{{.Result}}{{end}} +{{define "GET /ok ValueOK()"}}{{.Result}}{{end}} +{{define "GET /function Function()"}}{{.Result}}{{end}} +-- server.go -- +package server + +type T struct{} + +func (T) Value() string { return "" } +func (T) ValueError() (int, error) { return 0, nil } +func (T) ValueOK() (bool, bool) { return false, false } + +func Function() string { return "" } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + ValueError() (int, error) + ValueOK() (bool, bool) + Value() string +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /error", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, int]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + var err error + td.result, err = receiver.ValueError() + if err != nil { + td.errList = append(td.errList, err) + td.errStatusCode = http.StatusInternalServerError + } + } + if err := templates.ExecuteTemplate(buf, "GET /error ValueError()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /function", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = Function() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /function Function()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /ok", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, bool]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + var ok bool + td.result, ok = receiver.ValueOK() + if !ok { + return + } + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /ok ValueOK()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /value", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Value() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /value Value()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) ValueError() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "error") +} + +func (routePaths TemplateRoutePaths) Function() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "function") +} + +func (routePaths TemplateRoutePaths) ValueOk() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "ok") +} + +func (routePaths TemplateRoutePaths) Value() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "value") +} diff --git a/internal/generate/testdata/route_without_call.txtar b/internal/generate/testdata/route_without_call.txtar new file mode 100644 index 00000000..891cf1eb --- /dev/null +++ b/internal/generate/testdata/route_without_call.txtar @@ -0,0 +1,165 @@ +A route with no call renders its template with no result. +-- index.gohtml -- +{{define "GET /{$}"}}

    Home

    {{end}} +{{define "GET /about 203"}}

    About

    {{end}} +-- server.go -- +package server +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /about", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, struct { + }]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if err := templates.ExecuteTemplate(buf, "GET /about 203", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, http.StatusNonAuthoritativeInfo) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /{$}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, struct { + }]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if err := templates.ExecuteTemplate(buf, "GET /{$}", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) ReadAbout() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "about") +} + +func (routePaths TemplateRoutePaths) ReadExact() string { + return "/" +} diff --git a/internal/generate/testdata/sse.txtar b/internal/generate/testdata/sse.txtar new file mode 100644 index 00000000..a4195871 --- /dev/null +++ b/internal/generate/testdata/sse.txtar @@ -0,0 +1,342 @@ +An sse route streams an event per callback, from the route template or +a same-named one, and reads lastEventID. +-- index.gohtml -- +{{define "GET /events sse(Events(ctx, lastEventID, execute, sseClock))"}}{{.Result}}{{end}} +{{define "sseClock"}}{{.Result}}{{end}} +-- server.go -- +package server + +import "context" + +type T struct{} + +func (T) Events(ctx context.Context, lastEventID int, execute func(string) error, sseClock func(int) error) error { + return nil +} +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "context" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "path" + "strconv" + "strings" + "sync" +) + +type RoutesReceiver interface { + Events(ctx context.Context, lastEventID int, execute func(string) error, sseClock func(int) error) error +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /events", func(response http.ResponseWriter, request *http.Request) { + defer func() { + _ = request.Body.Close() + }() + flusher, ok := response.(http.Flusher) + if !ok { + http.Error(response, "streaming unsupported", http.StatusInternalServerError) + return + } + ctx := request.Context() + lastEventIDParsed, err := strconv.Atoi(request.Header.Get("Last-Event-Id")) + if err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + lastEventID := lastEventIDParsed + h := response.Header() + h.Set("Content-Type", "text/event-stream") + h.Set("Connection", "keep-alive") + h.Set("Cache-Control", "no-store") + response.WriteHeader(http.StatusOK) + flusher.Flush() + var mut sync.Mutex + if err := receiver.Events(ctx, lastEventID, func(result string) error { + if err := request.Context().Err(); err != nil { + return err + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + td := SSETemplateData[RoutesReceiver, string]{receiver: receiver, request: request, pathsPrefix: pathsPrefix, result: result} + if err := templates.ExecuteTemplate(buf, "GET /events sse(Events(ctx, lastEventID, execute, sseClock))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + return err + } + td.data = buf + mut.Lock() + defer mut.Unlock() + if _, err := td.WriteTo(response); err != nil { + return err + } + flusher.Flush() + return nil + }, func(result int) error { + if err := request.Context().Err(); err != nil { + return err + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + td := SSETemplateData[RoutesReceiver, int]{receiver: receiver, request: request, pathsPrefix: pathsPrefix, result: result} + if err := templates.ExecuteTemplate(buf, "sseClock", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + return err + } + td.data = buf + mut.Lock() + defer mut.Unlock() + if _, err := td.WriteTo(response); err != nil { + return err + } + flusher.Flush() + return nil + }); err != nil { + slog.ErrorContext(request.Context(), "sse handler returned an error", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + } + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type SSETemplateData[R, T any] struct { + receiver R + request *http.Request + result T + pathsPrefix string + event, id *string + retryMilliseconds *int + errList []error + data *bytes.Buffer +} + +func (m *SSETemplateData[R, T]) String() string { + return "" +} + +func (m *SSETemplateData[R, T]) Receiver() R { + return m.receiver +} + +func (m *SSETemplateData[R, T]) Request() *http.Request { + return m.request +} + +func (m *SSETemplateData[R, T]) Result() T { + return m.result +} + +func (m *SSETemplateData[R, T]) Err() error { + return errors.Join(m.errList...) +} + +func (m *SSETemplateData[R, T]) Event(event string) *SSETemplateData[R, T] { + m.event = &event + return m +} + +func (m *SSETemplateData[R, T]) ID(id string) *SSETemplateData[R, T] { + m.id = &id + return m +} + +func (m *SSETemplateData[R, T]) Retry(retryMilliseconds int) *SSETemplateData[R, T] { + m.retryMilliseconds = &retryMilliseconds + return m +} + +func (m *SSETemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: m.pathsPrefix} +} + +func (m *SSETemplateData[R, T]) WriteTo(w io.Writer) (int64, error) { + if m.id != nil && strings.ContainsAny(*m.id, "\r\n\x00") { + return 0, errors.New("sse: id contains a forbidden character") + } + if m.event != nil && strings.ContainsAny(*m.event, "\r\n") { + return 0, errors.New("sse: event contains a forbidden character") + } + var bytesWritten int + if m.id != nil { + if n, err := io.WriteString(w, "id: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := io.WriteString(w, *m.id); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.event != nil { + if n, err := io.WriteString(w, "event: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := io.WriteString(w, *m.event); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.retryMilliseconds != nil { + if n, err := io.WriteString(w, "retry: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + var retryBuf [20]byte + if n, err := w.Write(strconv.AppendInt(retryBuf[:0], int64(*m.retryMilliseconds), 10)); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + data := m.data.Bytes() + if bytes.IndexByte(data, '\r') >= 0 { + data = bytes.ReplaceAll(data, []byte("\r\n"), []byte("\n")) + data = bytes.ReplaceAll(data, []byte("\r"), []byte("\n")) + } + data = bytes.TrimSuffix(data, []byte{'\n'}) + for line := range bytes.SplitSeq(data, []byte{'\n'}) { + if n, err := io.WriteString(w, "data: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if len(line) > 0 { + if n, err := w.Write(line); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + return int64(bytesWritten), nil +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Events() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "events") +} diff --git a/internal/generate/testdata/sse_datastar.txtar b/internal/generate/testdata/sse_datastar.txtar new file mode 100644 index 00000000..6cc96367 --- /dev/null +++ b/internal/generate/testdata/sse_datastar.txtar @@ -0,0 +1,378 @@ +With --output-datastar an sse route frames datastar patch events, and +signals decodes the request body. +-- index.gohtml -- +{{define "POST /count sse(Count(signals, execute, countsSignals))"}}{{.Result}}{{end}} +-- server.go -- +package server + +type Signals struct{ Count int } + +type T struct{} + +func (T) Count(in Signals, execute func(int) error, countsSignals func(Signals) error) {} +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "path" + "strconv" + "strings" + "sync" +) + +type RoutesReceiver interface { + Count(in Signals, execute func(int) error, countsSignals func(Signals) error) +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("POST /count", func(response http.ResponseWriter, request *http.Request) { + defer func() { + _ = request.Body.Close() + }() + flusher, ok := response.(http.Flusher) + if !ok { + http.Error(response, "streaming unsupported", http.StatusInternalServerError) + return + } + var bodyValue Signals + if err := json.NewDecoder(request.Body).Decode(&bodyValue); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + h := response.Header() + h.Set("Content-Type", "text/event-stream") + h.Set("Connection", "keep-alive") + h.Set("Cache-Control", "no-store") + response.WriteHeader(http.StatusOK) + flusher.Flush() + var mut sync.Mutex + receiver.Count(bodyValue, func(result int) error { + if err := request.Context().Err(); err != nil { + return err + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + td := SSETemplateData[RoutesReceiver, int]{receiver: receiver, request: request, pathsPrefix: pathsPrefix, result: result} + if err := templates.ExecuteTemplate(buf, "POST /count sse(Count(signals, execute, countsSignals))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + return err + } + td.data = buf + mut.Lock() + defer mut.Unlock() + if _, err := td.WriteTo(response); err != nil { + return err + } + flusher.Flush() + return nil + }, func(result Signals) error { + if err := request.Context().Err(); err != nil { + return err + } + payload, err := json.Marshal(result) + if err != nil { + return err + } + mut.Lock() + defer mut.Unlock() + if _, err := io.WriteString(response, "event: datastar-patch-signals\ndata: signals "); err != nil { + return err + } + if _, err := response.Write(payload); err != nil { + return err + } + if _, err := io.WriteString(response, "\n\n"); err != nil { + return err + } + flusher.Flush() + return nil + }) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type SSETemplateData[R, T any] struct { + receiver R + request *http.Request + result T + pathsPrefix string + event, id *string + retryMilliseconds *int + errList []error + data *bytes.Buffer + selector, mode *string + useViewTransition bool +} + +func (m *SSETemplateData[R, T]) String() string { + return "" +} + +func (m *SSETemplateData[R, T]) Receiver() R { + return m.receiver +} + +func (m *SSETemplateData[R, T]) Request() *http.Request { + return m.request +} + +func (m *SSETemplateData[R, T]) Result() T { + return m.result +} + +func (m *SSETemplateData[R, T]) Err() error { + return errors.Join(m.errList...) +} + +func (m *SSETemplateData[R, T]) ID(id string) *SSETemplateData[R, T] { + m.id = &id + return m +} + +func (m *SSETemplateData[R, T]) Retry(retryMilliseconds int) *SSETemplateData[R, T] { + m.retryMilliseconds = &retryMilliseconds + return m +} + +func (m *SSETemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: m.pathsPrefix} +} + +func (m *SSETemplateData[R, T]) Selector(selector string) *SSETemplateData[R, T] { + m.selector = &selector + return m +} + +func (m *SSETemplateData[R, T]) Mode(mode string) *SSETemplateData[R, T] { + m.mode = &mode + return m +} + +func (m *SSETemplateData[R, T]) UseViewTransition(value bool) *SSETemplateData[R, T] { + m.useViewTransition = value + return m +} + +func (m *SSETemplateData[R, T]) WriteTo(w io.Writer) (int64, error) { + if m.id != nil && strings.ContainsAny(*m.id, "\r\n\x00") { + return 0, errors.New("sse: id contains a forbidden character") + } + if m.selector != nil && strings.ContainsAny(*m.selector, "\r\n") { + return 0, errors.New("sse: selector contains a forbidden character") + } + if m.mode != nil && strings.ContainsAny(*m.mode, "\r\n") { + return 0, errors.New("sse: mode contains a forbidden character") + } + var bytesWritten int + if n, err := io.WriteString(w, "event: datastar-patch-elements\n"); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if m.id != nil { + if n, err := io.WriteString(w, "id: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := io.WriteString(w, *m.id); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.retryMilliseconds != nil { + if n, err := io.WriteString(w, "retry: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + var retryBuf [20]byte + if n, err := w.Write(strconv.AppendInt(retryBuf[:0], int64(*m.retryMilliseconds), 10)); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.selector != nil { + if n, err := io.WriteString(w, "data: selector "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := io.WriteString(w, *m.selector); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.mode != nil { + if n, err := io.WriteString(w, "data: mode "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := io.WriteString(w, *m.mode); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.useViewTransition { + if n, err := io.WriteString(w, "data: useViewTransition true\n"); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + data := m.data.Bytes() + if bytes.IndexByte(data, '\r') >= 0 { + data = bytes.ReplaceAll(data, []byte("\r\n"), []byte("\n")) + data = bytes.ReplaceAll(data, []byte("\r"), []byte("\n")) + } + data = bytes.TrimSuffix(data, []byte{'\n'}) + for line := range bytes.SplitSeq(data, []byte{'\n'}) { + if n, err := io.WriteString(w, "data: elements "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write(line); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + return int64(bytesWritten), nil +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Count() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "count") +} diff --git a/internal/generate/testdata/sse_messages_without_datastar.txtar b/internal/generate/testdata/sse_messages_without_datastar.txtar new file mode 100644 index 00000000..c8db72ec --- /dev/null +++ b/internal/generate/testdata/sse_messages_without_datastar.txtar @@ -0,0 +1,314 @@ +An sse method returning nothing, with a data-less callback. +-- index.gohtml -- +{{define "GET /ticks/{id} sse(Ticks(request, id, execute))"}}tick{{end}} +-- server.go -- +package server + +import "net/http" + +type T struct{} + +func (T) Ticks(r *http.Request, id int, execute func() error) {} +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "path" + "strconv" + "strings" + "sync" +) + +type RoutesReceiver interface { + Ticks(r *http.Request, id int, execute func() error) +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /ticks/{id}", func(response http.ResponseWriter, request *http.Request) { + defer func() { + _ = request.Body.Close() + }() + flusher, ok := response.(http.Flusher) + if !ok { + http.Error(response, "streaming unsupported", http.StatusInternalServerError) + return + } + idParsed, err := strconv.Atoi(request.PathValue("id")) + if err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + idPathParam := idParsed + h := response.Header() + h.Set("Content-Type", "text/event-stream") + h.Set("Connection", "keep-alive") + h.Set("Cache-Control", "no-store") + response.WriteHeader(http.StatusOK) + flusher.Flush() + var mut sync.Mutex + receiver.Ticks(request, idPathParam, func() error { + if err := request.Context().Err(); err != nil { + return err + } + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + td := SSETemplateData[RoutesReceiver, struct{}]{receiver: receiver, request: request, pathsPrefix: pathsPrefix} + if err := templates.ExecuteTemplate(buf, "GET /ticks/{id} sse(Ticks(request, id, execute))", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + return err + } + td.data = buf + mut.Lock() + defer mut.Unlock() + if _, err := td.WriteTo(response); err != nil { + return err + } + flusher.Flush() + return nil + }) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type SSETemplateData[R, T any] struct { + receiver R + request *http.Request + result T + pathsPrefix string + event, id *string + retryMilliseconds *int + errList []error + data *bytes.Buffer +} + +func (m *SSETemplateData[R, T]) String() string { + return "" +} + +func (m *SSETemplateData[R, T]) Receiver() R { + return m.receiver +} + +func (m *SSETemplateData[R, T]) Request() *http.Request { + return m.request +} + +func (m *SSETemplateData[R, T]) Result() T { + return m.result +} + +func (m *SSETemplateData[R, T]) Err() error { + return errors.Join(m.errList...) +} + +func (m *SSETemplateData[R, T]) Event(event string) *SSETemplateData[R, T] { + m.event = &event + return m +} + +func (m *SSETemplateData[R, T]) ID(id string) *SSETemplateData[R, T] { + m.id = &id + return m +} + +func (m *SSETemplateData[R, T]) Retry(retryMilliseconds int) *SSETemplateData[R, T] { + m.retryMilliseconds = &retryMilliseconds + return m +} + +func (m *SSETemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: m.pathsPrefix} +} + +func (m *SSETemplateData[R, T]) WriteTo(w io.Writer) (int64, error) { + if m.id != nil && strings.ContainsAny(*m.id, "\r\n\x00") { + return 0, errors.New("sse: id contains a forbidden character") + } + if m.event != nil && strings.ContainsAny(*m.event, "\r\n") { + return 0, errors.New("sse: event contains a forbidden character") + } + var bytesWritten int + if m.id != nil { + if n, err := io.WriteString(w, "id: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := io.WriteString(w, *m.id); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.event != nil { + if n, err := io.WriteString(w, "event: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := io.WriteString(w, *m.event); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if m.retryMilliseconds != nil { + if n, err := io.WriteString(w, "retry: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + var retryBuf [20]byte + if n, err := w.Write(strconv.AppendInt(retryBuf[:0], int64(*m.retryMilliseconds), 10)); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + data := m.data.Bytes() + if bytes.IndexByte(data, '\r') >= 0 { + data = bytes.ReplaceAll(data, []byte("\r\n"), []byte("\n")) + data = bytes.ReplaceAll(data, []byte("\r"), []byte("\n")) + } + data = bytes.TrimSuffix(data, []byte{'\n'}) + for line := range bytes.SplitSeq(data, []byte{'\n'}) { + if n, err := io.WriteString(w, "data: "); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + if len(line) > 0 { + if n, err := w.Write(line); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + } + if n, err := w.Write([]byte{'\n'}); err != nil { + return int64(bytesWritten + n), err + } else { + bytesWritten += n + } + return int64(bytesWritten), nil +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Ticks(idPathParam int) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "ticks", strconv.Itoa(idPathParam)) +} diff --git a/internal/generate/testdata/status_codes.txtar b/internal/generate/testdata/status_codes.txtar new file mode 100644 index 00000000..bb470def --- /dev/null +++ b/internal/generate/testdata/status_codes.txtar @@ -0,0 +1,247 @@ +The status a route writes: one named in the template name, one the +result reports with a StatusCode method or field, or 200. +-- index.gohtml -- +{{define "POST /created 201 Create()"}}{{end}} +{{define "GET /constant http.StatusTeapot Create()"}}{{end}} +{{define "GET /method Method()"}}{{end}} +{{define "GET /field Field()"}}{{end}} +-- server.go -- +package server + +type T struct{} + +func (T) Create() string { return "" } + +type Coded struct{ code int } + +func (c Coded) StatusCode() int { return c.code } + +func (T) Method() Coded { return Coded{} } + +type WithField struct{ StatusCode int } + +func (T) Field() WithField { return WithField{} } +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "errors" + "fmt" + "log/slog" + "net/http" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Create() string + Field() WithField + Method() Coded +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /constant", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Create() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /constant http.StatusTeapot Create()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, http.StatusTeapot) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("POST /created", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Create() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "POST /created 201 Create()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, http.StatusCreated) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /field", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, WithField]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Field() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /field Field()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, td.result.StatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /method", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, Coded]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Method() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /method Method()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, td.result.StatusCode(), defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) ReadConstantCallingCreate() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "constant") +} + +func (routePaths TemplateRoutePaths) CreateCreatedCallingCreate() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "created") +} + +func (routePaths TemplateRoutePaths) Field() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "field") +} + +func (routePaths TemplateRoutePaths) Method() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "method") +} diff --git a/internal/generate/testdata/synthesized_method_note.txtar b/internal/generate/testdata/synthesized_method_note.txtar new file mode 100644 index 00000000..15f83c4d --- /dev/null +++ b/internal/generate/testdata/synthesized_method_note.txtar @@ -0,0 +1,199 @@ +With --use-receiver-type, a method the receiver does not define is +inferred, and generation says so. +-- index.gohtml -- +{{define "GET /{id} Missing(ctx, id)"}}{{.Result}}{{end}} +{{define "GET /defined Defined()"}}{{.Result}}{{end}} +-- server.go -- +package server + +type T struct{} + +func (T) Defined() string { return "" } +-- want/log.txt -- +note: T does not define Missing(ctx context.Context, id string) any +note: the inferred signatures return any — implement the methods to type-check the templates against real types +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "context" + "errors" + "fmt" + "log/slog" + "net/http" + "net/url" + "path" + "strconv" + "sync" +) + +type RoutesReceiver interface { + Defined() string + Missing(ctx context.Context, id string) any +} + +func TemplateRoutes(mux *http.ServeMux, receiver RoutesReceiver) TemplateRoutePaths { + pathsPrefix := "" + bytesBufferPool := sync.Pool{New: func() any { + return bytes.NewBuffer(nil) + }} + mux.HandleFunc("GET /defined", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, string]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Defined() + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /defined Defined()", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + mux.HandleFunc("GET /{id}", func(response http.ResponseWriter, request *http.Request) { + var td = TemplateData[RoutesReceiver, any]{receiver: receiver, response: response, request: request, pathsPrefix: pathsPrefix} + ctx := request.Context() + idPathParam := request.PathValue("id") + buf := bytesBufferPool.Get().(*bytes.Buffer) + buf.Reset() + defer bytesBufferPool.Put(buf) + if len(td.errList) == 0 { + td.result = receiver.Missing(ctx, idPathParam) + td.okay = true + } + if err := templates.ExecuteTemplate(buf, "GET /{id} Missing(ctx, id)", &td); err != nil { + slog.ErrorContext(request.Context(), "failed to render page", slog.String("path", request.URL.Path), slog.String("pattern", request.Pattern), slog.String("error", err.Error())) + http.Error(response, "failed to render page", http.StatusInternalServerError) + return + } + defaultStatusCode := http.StatusOK + if buf.Len() == 0 { + defaultStatusCode = http.StatusNoContent + } + statusCode := cmp.Or(td.statusCode, td.errStatusCode, defaultStatusCode) + if contentType := response.Header().Get("content-type"); contentType == "" { + response.Header().Set("content-type", "text/html; charset=utf-8") + } + response.Header().Set("content-length", strconv.Itoa(buf.Len())) + response.WriteHeader(statusCode) + _, _ = buf.WriteTo(response) + }) + return TemplateRoutePaths{pathsPrefix: pathsPrefix} +} + +type TemplateData[R any, T any] struct { + receiver R + response http.ResponseWriter + request *http.Request + result T + statusCode int + errStatusCode int + okay bool + errList []error + redirectURL string + pathsPrefix string +} + +func (data *TemplateData[R, T]) MuxtVersion() string { + const muxtVersion = "v1.2.3" + return muxtVersion +} + +func (data *TemplateData[R, T]) Path() TemplateRoutePaths { + return TemplateRoutePaths{pathsPrefix: data.pathsPrefix} +} + +func (data *TemplateData[R, T]) Result() T { + return data.result +} + +func (data *TemplateData[R, T]) Request() *http.Request { + return data.request +} + +func (data *TemplateData[R, T]) StatusCode(statusCode int) *TemplateData[R, T] { + data.statusCode = statusCode + return data +} + +func (data *TemplateData[R, T]) Header(key, value string) *TemplateData[R, T] { + data.response.Header().Set(key, value) + return data +} + +func (data *TemplateData[R, T]) Ok() bool { + return data.okay +} + +func (data *TemplateData[R, T]) Err() error { + return errors.Join(data.errList...) +} + +func (data *TemplateData[R, T]) Receiver() R { + return data.receiver +} + +func (data *TemplateData[R, T]) Redirect(url string, code int) (*TemplateData[R, T], error) { + if code < 300 || code >= 400 { + return data, fmt.Errorf("invalid status code %d for redirect", code) + } + data.redirectURL = url + return data.StatusCode(code), nil +} + +func (data *TemplateData[R, T]) RedirectMultipleChoices(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMultipleChoices) +} + +func (data *TemplateData[R, T]) RedirectMovedPermanently(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusMovedPermanently) +} + +func (data *TemplateData[R, T]) RedirectFound(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusFound) +} + +func (data *TemplateData[R, T]) RedirectSeeOther(url string) (*TemplateData[R, T], error) { + return data.Redirect(url, http.StatusSeeOther) +} + +func (data *TemplateData[R, T]) String() string { + return "" +} + +type TemplateRoutePaths struct { + pathsPrefix string +} + +func (routePaths TemplateRoutePaths) Defined() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "defined") +} + +func (routePaths TemplateRoutePaths) Missing(idPathParam string) string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), routePaths.escapePathSegment(idPathParam)) +} + +func (routePaths TemplateRoutePaths) escapePathSegment(value string) string { + switch value { + case ".": + return "%2E" + case "..": + return "%2E%2E" + } + return url.PathEscape(value) +} diff --git a/internal/generate/validation_test.go b/internal/generate/validation_test.go index 70516dc0..83009e29 100644 --- a/internal/generate/validation_test.go +++ b/internal/generate/validation_test.go @@ -5,7 +5,6 @@ import ( "go/ast" "go/types" "html/template" - "path/filepath" "strings" "testing" @@ -301,14 +300,7 @@ func Test_inputValidations(t *testing.T) { require.NoError(t, err) fragment := dom.NewDocumentFragment(nodes) - pl, err := loadPkg() - require.NoError(t, err) - fSet := fileSet() - wd, err := workingDir() - require.NoError(t, err) - - file, err := newFile(filepath.Join(wd, "tr.go"), fSet, pl) - require.NoError(t, err) + file := outputFile() input := fragment.QuerySelector(`[name="field"]`) require.NotNil(t, input) diff --git a/internal/asteval/diagnostic.go b/internal/load/diagnostic.go similarity index 99% rename from internal/asteval/diagnostic.go rename to internal/load/diagnostic.go index 8956b63c..8af09b6a 100644 --- a/internal/asteval/diagnostic.go +++ b/internal/load/diagnostic.go @@ -1,4 +1,4 @@ -package asteval +package load import ( "fmt" diff --git a/internal/asteval/diagnostic_test.go b/internal/load/diagnostic_test.go similarity index 90% rename from internal/asteval/diagnostic_test.go rename to internal/load/diagnostic_test.go index 3457a1bf..ff40476e 100644 --- a/internal/asteval/diagnostic_test.go +++ b/internal/load/diagnostic_test.go @@ -1,4 +1,4 @@ -package asteval_test +package load_test import ( "os" @@ -10,14 +10,14 @@ import ( "github.com/stretchr/testify/require" "golang.org/x/tools/go/packages" - "github.com/typelate/muxt/internal/asteval" + "github.com/typelate/muxt/internal/load" ) func TestNoPackageError(t *testing.T) { t.Run("no packages loaded", func(t *testing.T) { t.Setenv("GOWORK", "off") dir := t.TempDir() - err := asteval.NoPackageError(dir, nil) + err := load.NoPackageError(dir, nil) require.Error(t, err) require.Equal(t, "no Go package found at "+dir, err.Error(), "the short form is a single line") assert.Contains(t, multiLine(t, err), "loaded no packages") @@ -31,7 +31,7 @@ func TestNoPackageError(t *testing.T) { {Pos: "main.go:25:2", Msg: "undefined: TemplateRoutes"}, }}, } - err := asteval.NoPackageError(t.TempDir(), pl) + err := load.NoPackageError(t.TempDir(), pl) require.Error(t, err) assert.Contains(t, multiLine(t, err), "loaded 2 packages: fmt, example.com/broken") assert.Contains(t, multiLine(t, err), "undefined: TemplateRoutes") @@ -42,7 +42,7 @@ func TestNoPackageError(t *testing.T) { for range 5 { pkg.Errors = append(pkg.Errors, packages.Error{Msg: "boom"}) } - err := asteval.NoPackageError(t.TempDir(), []*packages.Package{pkg}) + err := load.NoPackageError(t.TempDir(), []*packages.Package{pkg}) require.Error(t, err) assert.Equal(t, 3, strings.Count(multiLine(t, err), "boom")) assert.Contains(t, multiLine(t, err), "more load errors omitted") @@ -53,7 +53,7 @@ func TestNoPackageError(t *testing.T) { require.NoError(t, os.WriteFile(filepath.Join(parent, "go.work"), []byte("go 1.24\n"), 0o600)) dir := filepath.Join(parent, "app") require.NoError(t, os.Mkdir(dir, 0o700)) - err := asteval.NoPackageError(dir, nil) + err := load.NoPackageError(dir, nil) require.Error(t, err) assert.Contains(t, multiLine(t, err), filepath.Join(parent, "go.work")) assert.Contains(t, multiLine(t, err), "GOWORK=off") @@ -65,14 +65,14 @@ func TestNoPackageError(t *testing.T) { require.NoError(t, os.WriteFile(filepath.Join(parent, "go.work"), []byte("go 1.24\n"), 0o600)) dir := filepath.Join(parent, "app") require.NoError(t, os.Mkdir(dir, 0o700)) - err := asteval.NoPackageError(dir, nil) + err := load.NoPackageError(dir, nil) require.Error(t, err) assert.Contains(t, multiLine(t, err), filepath.Join(parent, "go.work")) assert.NotContains(t, multiLine(t, err), "GOWORK=auto is set") }) t.Run("an explicit GOWORK is named", func(t *testing.T) { t.Setenv("GOWORK", "/somewhere/go.work") - err := asteval.NoPackageError(t.TempDir(), nil) + err := load.NoPackageError(t.TempDir(), nil) require.Error(t, err) assert.Contains(t, multiLine(t, err), "GOWORK=/somewhere/go.work is set") }) @@ -84,7 +84,7 @@ func TestNoPackageError(t *testing.T) { pkgDir := filepath.Join(module, "internal", "hypertext") require.NoError(t, os.MkdirAll(pkgDir, 0o700)) require.NoError(t, os.WriteFile(filepath.Join(module, "go.mod"), []byte("module app\n\ngo 1.24\n"), 0o600)) - err := asteval.NoPackageError(pkgDir, nil) + err := load.NoPackageError(pkgDir, nil) require.Error(t, err) assert.Contains(t, multiLine(t, err), "go work use "+module) assert.NotContains(t, multiLine(t, err), "go work use "+pkgDir) @@ -93,7 +93,7 @@ func TestNoPackageError(t *testing.T) { t.Setenv("GOWORK", "off") dir := t.TempDir() pl := []*packages.Package{{PkgPath: dir, Errors: []packages.Error{{Msg: "contained in a module that is not one of the workspace modules"}}}} - err := asteval.NoPackageError(dir, pl) + err := load.NoPackageError(dir, pl) require.Error(t, err) assert.Equal(t, "the Go package at "+dir+" loaded, but with errors", err.Error()) assert.NotContains(t, multiLine(t, err), "no Go package found") @@ -105,7 +105,7 @@ func TestNoPackageError(t *testing.T) { require.NoError(t, os.WriteFile(filepath.Join(dir, "go.mod"), []byte("module broken\n\ngo 1.24\n"), 0o600)) require.NoError(t, os.WriteFile(filepath.Join(dir, "go.work"), []byte("go 1.24\n\nuse (\n\t.\n\t./missing\n)\n"), 0o600)) t.Setenv("GOWORK", filepath.Join(dir, "go.work")) - _, _, err := asteval.LoadPackages(dir) + _, _, err := load.Packages(dir) require.Error(t, err) assert.Contains(t, err.Error(), "failed to load Go packages from "+dir) msg := multiLine(t, err) @@ -115,7 +115,7 @@ func TestNoPackageError(t *testing.T) { }) t.Run("GOWORK off adds no workspace note", func(t *testing.T) { t.Setenv("GOWORK", "off") - err := asteval.NoPackageError(t.TempDir(), nil) + err := load.NoPackageError(t.TempDir(), nil) require.Error(t, err) assert.NotContains(t, multiLine(t, err), "go work use") }) diff --git a/internal/load/hydrate.go b/internal/load/hydrate.go new file mode 100644 index 00000000..26ca0b66 --- /dev/null +++ b/internal/load/hydrate.go @@ -0,0 +1,52 @@ +package load + +import ( + "go/types" + "path/filepath" + + "golang.org/x/tools/go/packages" + + "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/generate" + "github.com/typelate/muxt/internal/source" +) + +// This file hydrates a command's configuration: it reads, from the packages +// loaded for the working directory, what the configuration names, and +// returns it as the input the command's implementation runs on. + +// GenerateSource hydrates muxt generate's configuration: the package the +// routes file is written into, with its templates variables, and the +// receiver type the configuration names, nil when it names none. +// +// The routes file belongs to the package in its own directory, which is +// the working directory unless the output file names another. +func GenerateSource(wd string, pl []*packages.Package, config generate.RoutesFileConfiguration) (source.Package, *types.Named, error) { + return packageWithReceiver(filepath.Dir(filepath.Join(wd, config.OutputFileName)), pl, config.ReceiverPackage, config.ReceiverType, config.TemplatesVariables) +} + +// RoutesSource hydrates the route listing's configuration. +func RoutesSource(wd string, pl []*packages.Package, config analysis.DefinitionsConfiguration) (source.Package, *types.Named, error) { + return packageWithReceiver(wd, pl, config.ReceiverPackage, config.ReceiverType, config.TemplatesVariables) +} + +// packageWithReceiver reads the package at dir and the receiver type, +// reporting a missing package before a missing receiver, and either before +// a templates variable that does not evaluate. +func packageWithReceiver(dir string, pl []*packages.Package, receiverPackage, receiverType string, variables []string) (source.Package, *types.Named, error) { + if _, ok := PackageAtFilepath(pl, dir); !ok { + return source.Package{}, nil, NoPackageError(dir, pl) + } + var receiver *types.Named + if receiverType != "" { + var err error + if receiver, err = Receiver(dir, pl, receiverPackage, receiverType); err != nil { + return source.Package{}, nil, err + } + } + pkg, err := Package(dir, pl, variables) + if err != nil { + return source.Package{}, nil, err + } + return pkg, receiver, nil +} diff --git a/internal/load/hydrate_test.go b/internal/load/hydrate_test.go new file mode 100644 index 00000000..bced323c --- /dev/null +++ b/internal/load/hydrate_test.go @@ -0,0 +1,82 @@ +package load_test + +import ( + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/generate" + "github.com/typelate/muxt/internal/load" + "github.com/typelate/muxt/internal/load/loadtest" +) + +const hydrateServer = `package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +type Server struct{} + +func Render(w io.Writer) error { + return templates.ExecuteTemplate(w, "page", nil) +} +` + +// TestHydration states what a command's configuration reads from a load, +// and which failure a configuration naming several missing things reports. +func TestHydration(t *testing.T) { + dir := t.TempDir() + pl := loadtest.Package(t, dir, "example.com/server", map[string]string{ + "server.go": hydrateServer, + "page.gohtml": `{{define "page"}}

    {{.}}

    {{end}}`, + }) + + t.Run("the package, its receiver, and each variable", func(t *testing.T) { + pkg, receiver, err := load.RoutesSource(dir, pl, analysis.DefinitionsConfiguration{ReceiverType: "Server", TemplatesVariables: []string{"templates"}}) + require.NoError(t, err) + assert.Equal(t, "Server", receiver.Obj().Name()) + require.Len(t, pkg.Variables, 1) + variable := pkg.Variables[0] + assert.Equal(t, "templates", variable.Name) + require.Len(t, variable.Calls, 1) + assert.Equal(t, "page", variable.Calls[0].Template) + assert.Equal(t, filepath.Join(dir, "server.go"), variable.Calls[0].Position.Filename) + _, ok := variable.NamePosition("page") + assert.True(t, ok, "the page's definition is located") + _, ok = pkg.Import("net/http") + assert.True(t, ok, "the standard library a load always includes is importable") + }) + + t.Run("no receiver named is none looked up", func(t *testing.T) { + _, receiver, err := load.RoutesSource(dir, pl, analysis.DefinitionsConfiguration{TemplatesVariables: []string{"templates"}}) + require.NoError(t, err) + assert.Nil(t, receiver) + }) + + t.Run("a missing receiver is reported before a missing variable", func(t *testing.T) { + _, _, err := load.RoutesSource(dir, pl, analysis.DefinitionsConfiguration{ReceiverType: "Srever", TemplatesVariables: []string{"nope"}}) + require.EqualError(t, err, "could not find receiver type Srever in example.com/server; did you mean Server?") + }) + + t.Run("the first variable that does not evaluate", func(t *testing.T) { + _, err := load.Package(dir, pl, []string{"templates", "nope", "neither"}) + require.EqualError(t, err, "variable nope not found in package example.com/server") + }) + + t.Run("the routes file belongs to the package in its own directory", func(t *testing.T) { + _, _, err := load.GenerateSource(dir, pl, generate.RoutesFileConfiguration{OutputFileName: filepath.Join("sub", "routes.go"), ReceiverType: "Srever"}) + require.Error(t, err) + assert.Equal(t, "no Go package found at "+filepath.Join(dir, "sub"), err.Error(), "a missing package is reported before a missing receiver") + }) +} diff --git a/internal/load/loadtest/loadtest.go b/internal/load/loadtest/loadtest.go new file mode 100644 index 00000000..7575f351 --- /dev/null +++ b/internal/load/loadtest/loadtest.go @@ -0,0 +1,86 @@ +// Package loadtest builds the result of a package load without running +// the go command, for tests of what internal/load and its callers do with +// a loaded package. +// +// packages.Load runs go list, which takes seconds. What muxt reads from +// its result -- syntax, type information, the embedded files -- is plain +// data, and internal/typestest produces the type information in +// microseconds. Package writes the files to disk, because evaluating +// ParseFS and reading template sources both read them from there, and +// assembles the *packages.Package go list would have reported. +package loadtest + +import ( + "os" + "path" + "path/filepath" + "slices" + "strings" + "testing" + + "golang.org/x/tools/go/packages" + + "github.com/typelate/muxt/internal/typestest" +) + +// Package writes files into dir and returns them loaded as the package +// with import path pkgPath, the way load.Packages would report a load of +// dir: that package, then the standard library packages every load +// includes. +// +// Go files are type checked against the stub standard library in +// internal/typestest. Every other file is embedded, as a //go:embed +// pattern covering it would. Keys are slash-separated paths relative to +// dir; files in subdirectories are written but not part of the package. +func Package(t testing.TB, dir, pkgPath string, files map[string]string) []*packages.Package { + t.Helper() + goFiles := make(map[string]string) + var goPaths, embedded []string + for name, content := range files { + file := filepath.Join(dir, filepath.FromSlash(name)) + if err := os.MkdirAll(filepath.Dir(file), 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(file, []byte(content), 0o600); err != nil { + t.Fatal(err) + } + switch { + case path.Dir(name) != ".": + case strings.HasSuffix(name, ".go"): + goFiles[file] = content + goPaths = append(goPaths, file) + default: + embedded = append(embedded, file) + } + } + slices.Sort(goPaths) + slices.Sort(embedded) + + checked, err := typestest.CheckSyntax(pkgPath, goFiles) + if err != nil { + t.Fatal(err) + } + pl := []*packages.Package{{ + ID: pkgPath, + Name: checked.Types.Name(), + PkgPath: pkgPath, + Dir: dir, + GoFiles: goPaths, + EmbedFiles: embedded, + Fset: typestest.FileSet, + Syntax: checked.Syntax, + Types: checked.Types, + TypesInfo: checked.Info, + }} + // load.Packages always loads these alongside the working directory's + // package, so what a route argument binds to is found even when the + // package does not import it. + for _, path := range []string{"encoding", "fmt", "net/http"} { + std, ok := typestest.Lookup(path) + if !ok { + t.Fatalf("typestest has no stub for %s", path) + } + pl = append(pl, &packages.Package{ID: path, Name: std.Name(), PkgPath: path, Fset: typestest.FileSet, Types: std}) + } + return pl +} diff --git a/internal/load/loadtest/loadtest_test.go b/internal/load/loadtest/loadtest_test.go new file mode 100644 index 00000000..166119b3 --- /dev/null +++ b/internal/load/loadtest/loadtest_test.go @@ -0,0 +1,64 @@ +package loadtest_test + +import ( + "path/filepath" + "slices" + "testing" + + "github.com/typelate/muxt/internal/load" + "github.com/typelate/muxt/internal/load/loadtest" +) + +func TestPackageLoadsTemplates(t *testing.T) { + dir := t.TempDir() + pl := loadtest.Package(t, dir, "example.com/server", map[string]string{ + "server.go": `package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +var inline = template.Must(template.New("inline").Delims("[[", "]]").Parse(` + "`[[define \"note\"]][[.]][[end]]`" + `)) + +func Render(w io.Writer, name string) error { + return templates.ExecuteTemplate(w, "page", name) +} +`, + "page.gohtml": `{{define "page"}}

    {{.}}

    {{end}}`, + }) + + pkg, err := load.Package(dir, pl, []string{"templates", "inline"}) + if err != nil { + t.Fatal(err) + } + if pkg.Types.Path() != "example.com/server" { + t.Errorf("package %s, want example.com/server", pkg.Types.Path()) + } + sets := pkg.Variables + + if sets[0].Set.Lookup("page") == nil { + t.Error("templates does not hold the page ParseFS read") + } + definition, ok := sets[0].Definitions["page"] + if !ok || definition.Define.Filename != filepath.Join(dir, "page.gohtml") { + t.Errorf("page defined at %+v, want in page.gohtml", definition.Define) + } + if got := len(sets[0].Calls); got != 1 { + t.Errorf("%d ExecuteTemplate calls, want the one Render makes", got) + } + var names []string + for _, tmpl := range sets[1].Set.Templates() { + names = append(names, tmpl.Name()) + } + slices.Sort(names) + if !slices.Equal(names, []string{"inline", "note"}) { + t.Errorf("inline holds %q, want inline and note", names) + } +} diff --git a/internal/asteval/package.go b/internal/load/package.go similarity index 74% rename from internal/asteval/package.go rename to internal/load/package.go index 226d51b2..63febaab 100644 --- a/internal/asteval/package.go +++ b/internal/load/package.go @@ -1,4 +1,4 @@ -package asteval +package load import ( "fmt" @@ -6,6 +6,7 @@ import ( "go/types" "html/template" "path/filepath" + "strings" "github.com/typelate/check" @@ -13,13 +14,13 @@ import ( "golang.org/x/tools/go/packages" ) -func LoadPackages(wd string, morePatterns ...string) (*token.FileSet, []*packages.Package, error) { - return LoadPackagesWithEnv(wd, nil, morePatterns...) +func Packages(wd string, morePatterns ...string) (*token.FileSet, []*packages.Package, error) { + return PackagesWithEnv(wd, nil, morePatterns...) } -// LoadPackagesWithEnv is LoadPackages with the environment the go command +// PackagesWithEnv is Packages with the environment the go command // runs in. A nil env is the process's own. -func LoadPackagesWithEnv(wd string, env []string, morePatterns ...string) (*token.FileSet, []*packages.Package, error) { +func PackagesWithEnv(wd string, env []string, morePatterns ...string) (*token.FileSet, []*packages.Package, error) { patterns := []string{ wd, "encoding", "fmt", "net/http", } @@ -41,6 +42,42 @@ func LoadPackagesWithEnv(wd string, env []string, morePatterns ...string) (*toke return fileSet, pl, err } +// PackagesWithTests is PackagesWithEnv, loading each package with its +// in-package test files too. +// +// With tests, go list reports a package twice: once as it is written and +// once compiled with its test files. The second holds a test's +// ExecuteTemplate calls, so the test variants come first, ahead of the +// packages as written, and a package picked by directory is the one with +// its tests. Unlike PackagesWithEnv, a load failure is returned as the +// loader reported it. +func PackagesWithTests(wd string, env []string) ([]*packages.Package, error) { + pl, err := packages.Load(&packages.Config{ + Fset: token.NewFileSet(), + Tests: true, + Mode: packages.NeedModule | packages.NeedTypesInfo | packages.NeedName | + packages.NeedFiles | packages.NeedTypes | packages.NeedSyntax | + packages.NeedEmbedPatterns | packages.NeedEmbedFiles | packages.NeedImports, + Dir: wd, + Env: env, + }, wd, "encoding", "fmt", "net/http") + if err != nil { + return nil, err + } + ordered := make([]*packages.Package, 0, len(pl)) + for _, pkg := range pl { + if strings.HasSuffix(pkg.ID, ".test]") { + ordered = append(ordered, pkg) + } + } + for _, pkg := range pl { + if !strings.HasSuffix(pkg.ID, ".test]") { + ordered = append(ordered, pkg) + } + } + return ordered, nil +} + // ParseErrors returns the syntax errors the loader recovered from. // // The loader carries on with a partial AST, so a caller still gets @@ -95,42 +132,6 @@ func PackageWithPath(list []*packages.Package, path string) (*packages.Package, return nil, false } -// LoadedTemplates bundles a package's loaded template variable with the -// analysis wiring built from it. -type LoadedTemplates struct { - Package *packages.Package - Templates *check.Templates - Global *check.Global - HTML *template.Template -} - -func LoadTemplates(wd, templatesVariable string, pl []*packages.Package) (*LoadedTemplates, error) { - pkg, ok := PackageAtFilepath(pl, wd) - if !ok { - return nil, NoPackageError(wd, pl) - } - - lt, ts, err := HTMLTemplates(templatesVariable, pkg) - if err != nil { - return nil, err - } - - global := check.NewGlobal(pkg.Types, pkg.Fset, lt, lt.Functions()) - global.Definitions = lt - return &LoadedTemplates{Package: pkg, Templates: lt, Global: global, HTML: ts}, nil -} - -// Templates evaluates the package-level template variable through -// check.LoadTemplates and returns the html/template value together with -// the functions collected from Funcs calls in its construction chain. -func Templates(templatesVariable string, pkg *packages.Package) (*template.Template, check.Functions, error) { - lt, ts, err := HTMLTemplates(templatesVariable, pkg) - if err != nil { - return nil, nil, err - } - return ts, lt.CollectedFunctions(), nil -} - // HTMLTemplates evaluates the package-level template variable through // check.LoadTemplates and returns the loaded handle alongside the // html/template value; muxt introspects template names and trees without diff --git a/internal/asteval/package_test.go b/internal/load/package_test.go similarity index 99% rename from internal/asteval/package_test.go rename to internal/load/package_test.go index cc5b1ee6..fe6a16d8 100644 --- a/internal/asteval/package_test.go +++ b/internal/load/package_test.go @@ -1,4 +1,4 @@ -package asteval +package load import ( "slices" diff --git a/internal/load/source.go b/internal/load/source.go new file mode 100644 index 00000000..c09b85ab --- /dev/null +++ b/internal/load/source.go @@ -0,0 +1,114 @@ +package load + +import ( + "cmp" + "go/types" + + "golang.org/x/tools/go/packages" + + "github.com/typelate/muxt/internal/source" +) + +// Package reads the package at dir among pl, with the templates variables +// named, in order, into a source.Package. It is where a go/packages result +// stops: past it, a run holds go/types values and template sets and +// nothing that needs the go command. +// +// It fails when no loaded package is at dir, or at the first variable that +// does not evaluate to a template set. +func Package(dir string, pl []*packages.Package, variables []string) (source.Package, error) { + pkg, ok := PackageAtFilepath(pl, dir) + if !ok { + return source.Package{}, NoPackageError(dir, pl) + } + result := source.Package{ + Fset: pkg.Fset, + Types: pkg.Types, + Imports: imports(pl), + } + for _, name := range variables { + variable, err := Variable(pkg, name) + if err != nil { + return source.Package{}, err + } + result.Variables = append(result.Variables, variable) + } + return result, nil +} + +// Variable evaluates the templates variable name in pkg: its template set, +// the functions its templates may call, where each template was defined, +// and the ExecuteTemplate calls made on it. +func Variable(pkg *packages.Package, name string) (source.Variable, error) { + lt, ts, err := HTMLTemplates(name, pkg) + if err != nil { + return source.Variable{}, err + } + variable := source.Variable{ + Name: name, + Set: ts, + Functions: lt.Functions(), + Funcs: lt.CollectedFunctions(), + Definitions: make(map[string]source.Definition), + } + for _, t := range ts.Templates() { + d, ok := lt.FindDefinition(t.Name()) + if !ok { + continue + } + variable.Definitions[t.Name()] = source.Definition{ + Name: d.Name, + Define: source.Span{Position: d.Define.Position, Length: d.Define.Length}, + End: source.Span{Position: d.End.Position, Length: d.End.Length}, + TemplateName: source.Span{Position: d.TemplateName.Position, Length: d.TemplateName.Length}, + Tree: d.Tree, + } + } + for call := range lt.ExecuteTemplateCalls() { + variable.Calls = append(variable.Calls, source.Call{ + Position: pkg.Fset.Position(call.Call.Pos()), + Template: call.TemplateName, + Data: call.DataType, + }) + } + return variable, nil +} + +// Receiver finds the receiver type named ident in the package at dir among +// pl, or in the package with import path packagePath when it is set. +func Receiver(dir string, pl []*packages.Package, packagePath, ident string) (*types.Named, error) { + pkg, ok := PackageAtFilepath(pl, dir) + if !ok { + return nil, NoPackageError(dir, pl) + } + return FindType(pl, cmp.Or(packagePath, pkg.PkgPath), ident) +} + +// imports indexes, by import path, every package in pl and every package +// they import. A path loaded more than once -- a package and its test +// variant -- keeps the first in pl. +func imports(pl []*packages.Package) map[string]*types.Package { + index := make(map[string]*types.Package) + var queue []*types.Package + for _, pkg := range pl { + if pkg.Types == nil { + continue + } + if _, seen := index[pkg.Types.Path()]; !seen { + index[pkg.Types.Path()] = pkg.Types + queue = append(queue, pkg.Types) + } + } + for len(queue) > 0 { + pkg := queue[0] + queue = queue[1:] + for _, imported := range pkg.Imports() { + if _, seen := index[imported.Path()]; seen { + continue + } + index[imported.Path()] = imported + queue = append(queue, imported) + } + } + return index +} diff --git a/internal/asteval/asteval_test.go b/internal/load/templates_test.go similarity index 78% rename from internal/asteval/asteval_test.go rename to internal/load/templates_test.go index c96c5fc9..37051874 100644 --- a/internal/asteval/asteval_test.go +++ b/internal/load/templates_test.go @@ -1,4 +1,4 @@ -package asteval_test +package load_test import ( "os" @@ -9,7 +9,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/typelate/muxt/internal/asteval" + "github.com/typelate/muxt/internal/load" ) // writeModule lays out a scratch module so templates load the way muxt @@ -51,14 +51,15 @@ func TestTemplates(t *testing.T) { "index.gohtml": `{{define "home"}}Hello, {{upper .Name}}{{end}}`, "form.gohtml": `{{define "create"}}
    {{end}}`, }) - _, pl, err := asteval.LoadPackages(dir) + _, pl, err := load.Packages(dir) require.NoError(t, err) - pkg, ok := asteval.PackageAtFilepath(pl, dir) + pkg, ok := load.PackageAtFilepath(pl, dir) require.True(t, ok) t.Run("parses the embedded files", func(t *testing.T) { - ts, functions, err := asteval.Templates("templates", pkg) + lt, ts, err := load.HTMLTemplates("templates", pkg) require.NoError(t, err) + functions := lt.CollectedFunctions() var names []string for _, tmpl := range ts.Templates() { @@ -73,19 +74,21 @@ func TestTemplates(t *testing.T) { }) t.Run("unknown variable", func(t *testing.T) { - _, _, err := asteval.Templates("nope", pkg) + _, _, err := load.HTMLTemplates("nope", pkg) require.ErrorContains(t, err, "variable nope not found") }) - t.Run("load templates wires the global", func(t *testing.T) { - lt, err := asteval.LoadTemplates(dir, "templates", pl) + t.Run("a variable locates its definitions and calls", func(t *testing.T) { + variable, err := load.Variable(pkg, "templates") require.NoError(t, err) - require.NotNil(t, lt.HTML) + require.NotNil(t, variable.Set) - def, ok := lt.Global.Definitions.FindDefinition("home") + def, ok := variable.Definitions["home"] require.True(t, ok, "definitions resolve for file-parsed templates") require.True(t, def.Define.IsValid()) assert.Equal(t, "index.gohtml", filepath.Base(def.Define.Filename)) + _, ok = variable.Funcs["upper"] + assert.True(t, ok, "Funcs holds the Funcs-registered functions") }) } @@ -100,14 +103,14 @@ var texts = template.Must(template.New("t").Parse(` + "`{{define \"note\"}}hi{{e func main() {} `, }) - _, pl, err := asteval.LoadPackages(dir) + _, pl, err := load.Packages(dir) require.NoError(t, err) - pkg, ok := asteval.PackageAtFilepath(pl, dir) + pkg, ok := load.PackageAtFilepath(pl, dir) require.True(t, ok) // Muxt introspects trees without executing, so a text/template // variable loads through an html/template value with the same trees. - ts, _, err := asteval.Templates("texts", pkg) + _, ts, err := load.HTMLTemplates("texts", pkg) require.NoError(t, err) require.NotNil(t, ts.Lookup("note")) } diff --git a/internal/mutation/collect.go b/internal/mutation/collect.go index 4a9ad14a..f09c15ff 100644 --- a/internal/mutation/collect.go +++ b/internal/mutation/collect.go @@ -6,8 +6,7 @@ import ( "path/filepath" "slices" - "github.com/typelate/check" - "golang.org/x/tools/go/packages" + "github.com/typelate/muxt/internal/source" ) // sourceKey identifies the text a template was written in: a template @@ -21,7 +20,6 @@ type sourceKey struct { // distinct texts holding them, reading each file once. type sourceCollector struct { workingDirectory string - packages []*packages.Package files map[string]string byKey map[sourceKey]*templateSource keys []sourceKey @@ -47,7 +45,7 @@ type sourceCollector struct { // Within one source the pair is fixed, so the first definition that // reveals it answers for the whole source. A source whose only template // has no define clause reveals nothing and keeps the defaults. -func (c *sourceCollector) resolveDelimiters(defs []check.Definition) { +func (c *sourceCollector) resolveDelimiters(defs []source.Definition) { for _, definition := range defs { file := definition.Define.Position.Filename if file == "" { @@ -74,22 +72,25 @@ func (c *sourceCollector) resolveDelimiters(defs []check.Definition) { // // It is the same key add files the definition under, so the delimiters // resolved here reach the source they were read from. -func (c *sourceCollector) keyFor(definition check.Definition) (sourceKey, bool) { +func (c *sourceCollector) keyFor(definition source.Definition) (sourceKey, bool) { file := definition.Define.Position.Filename if filepath.Ext(file) != ".go" { return sourceKey{file: file}, true } - litStart, _, ok := findStringLiteral(c.packages, file, definition.Define.Offset) + text, err := c.read(file) + if err != nil { + return sourceKey{}, false + } + litStart, _, ok := findStringLiteral(file, text, definition.Define.Offset) if !ok { return sourceKey{}, false } return sourceKey{file: file, litStart: litStart}, true } -func newSourceCollector(workingDirectory string, pl []*packages.Package, defs []check.Definition) *sourceCollector { +func newSourceCollector(workingDirectory string, defs []source.Definition) *sourceCollector { c := &sourceCollector{ workingDirectory: workingDirectory, - packages: pl, files: make(map[string]string), byKey: make(map[sourceKey]*templateSource), delims: make(map[sourceKey][2]string), @@ -108,10 +109,10 @@ func (c *sourceCollector) delimitersFor(key sourceKey) (string, string) { // add files one definition under the text it was written in, reading // that text at most once, and reports the source it was filed under. // -// A definition the collector cannot place -- a Go string literal it has -// no package for -- is reported as no source rather than as an error, +// A definition the collector cannot place -- an offset in a Go file that +// is not inside a string literal -- is reported as no source rather than as an error, // since the caller may hold others it can still use. -func (c *sourceCollector) add(definition check.Definition) (*templateSource, error) { +func (c *sourceCollector) add(definition source.Definition) (*templateSource, error) { file := definition.Define.Position.Filename if file == "" { return nil, nil @@ -131,7 +132,7 @@ func (c *sourceCollector) add(definition check.Definition) (*templateSource, err return nil, err } } else { - litStart, litEnd, ok := findStringLiteral(c.packages, file, definition.Define.Offset) + litStart, litEnd, ok := findStringLiteral(file, fileText, definition.Define.Offset) if !ok { return nil, nil } diff --git a/internal/mutation/delimiters_test.go b/internal/mutation/delimiters_test.go index 26fdabe1..c42a7d2a 100644 --- a/internal/mutation/delimiters_test.go +++ b/internal/mutation/delimiters_test.go @@ -5,7 +5,7 @@ import ( "strings" "testing" - "github.com/typelate/check" + "github.com/typelate/muxt/internal/source" ) // TestDelimitersReadsThemOffTheEndClause states how the delimiters a @@ -51,11 +51,11 @@ func TestDelimitersReadsThemOffTheEndClause(t *testing.T) { // trivially at offset zero. const prefix = "hello " text := prefix + tt.end - definition := check.Definition{ + definition := source.Definition{ Name: "x", - Define: check.Span{Position: token.Position{Filename: "t.gohtml"}, Length: 1}, - TemplateName: check.Span{Position: token.Position{Filename: "t.gohtml", Offset: 1, Line: 1}, Length: 3}, - End: check.Span{ + Define: source.Span{Position: token.Position{Filename: "t.gohtml"}, Length: 1}, + TemplateName: source.Span{Position: token.Position{Filename: "t.gohtml", Offset: 1, Line: 1}, Length: 3}, + End: source.Span{ Position: token.Position{Filename: "t.gohtml", Offset: len(prefix)}, Length: len(tt.end), }, @@ -82,9 +82,9 @@ func TestDelimitersDeclinesWhatItCannotRead(t *testing.T) { const text = `{{define "x"}}{{end}}` t.Run("a template with no define clause", func(t *testing.T) { - definition := check.Definition{ + definition := source.Definition{ Name: "t.gohtml", - End: check.Span{Position: token.Position{Filename: "t.gohtml", Offset: len(text)}}, + End: source.Span{Position: token.Position{Filename: "t.gohtml", Offset: len(text)}}, } if _, _, ok := delimiters(text, definition); ok { t.Error("delimiters accepted a definition with no end clause to read") @@ -92,10 +92,10 @@ func TestDelimitersDeclinesWhatItCannotRead(t *testing.T) { }) t.Run("a span outside the text", func(t *testing.T) { - definition := check.Definition{ + definition := source.Definition{ Name: "x", - TemplateName: check.Span{Position: token.Position{Filename: "t.gohtml", Offset: 9, Line: 1}, Length: 3}, - End: check.Span{ + TemplateName: source.Span{Position: token.Position{Filename: "t.gohtml", Offset: 9, Line: 1}, Length: 3}, + End: source.Span{ Position: token.Position{Filename: "t.gohtml", Offset: len(text)}, Length: 99, }, @@ -116,11 +116,11 @@ func TestDelimitersAgreeWithTheScanner(t *testing.T) { const text = `[[define "greeting"]]Hello, [[.Name]]![[end]]` endAt := strings.LastIndex(text, "[[end]]") - definition := check.Definition{ + definition := source.Definition{ Name: "greeting", - Define: check.Span{Position: token.Position{Filename: "t.gohtml"}, Length: 21}, - TemplateName: check.Span{Position: token.Position{Filename: "t.gohtml", Offset: 9, Line: 1}, Length: 10}, - End: check.Span{Position: token.Position{Filename: "t.gohtml", Offset: endAt}, Length: len("[[end]]")}, + Define: source.Span{Position: token.Position{Filename: "t.gohtml"}, Length: 21}, + TemplateName: source.Span{Position: token.Position{Filename: "t.gohtml", Offset: 9, Line: 1}, Length: 10}, + End: source.Span{Position: token.Position{Filename: "t.gohtml", Offset: endAt}, Length: len("[[end]]")}, } left, right, ok := delimiters(text, definition) diff --git a/internal/mutation/diff.go b/internal/mutation/diff.go index 2185ce9e..79d13f40 100644 --- a/internal/mutation/diff.go +++ b/internal/mutation/diff.go @@ -11,8 +11,6 @@ import ( "os/exec" "path/filepath" "strings" - - "github.com/typelate/muxt/internal/asteval" ) // revision is what the templates looked like at an earlier commit: the @@ -56,17 +54,20 @@ func scopesOf(scopes []scope) revision { func templatesAt(config Configuration, dir string) (revision, error) { // The copy is outside any workspace GOWORK may name, and would fail to // load within one, so it loads as the module it is. - pl, err := loadPackages(dir, config.IncludeTests, append(config.environment(), "GOWORK=off")) + in, err := loadInput(dir, config, append(config.environment(), "GOWORK=off")) if err != nil { return nil, err } + return revisionOf(in) +} + +// revisionOf records what each template reached in the input reads like, +// which is what a --diff run compares the working tree with. +func revisionOf(in input) (revision, error) { before := make(revision) - for _, templatesVariable := range config.TemplatesVariables { - lt, err := asteval.LoadTemplates(dir, templatesVariable, pl) - if err != nil { - return nil, err - } - index, err := buildTreeIndex(lt, dir, pl, lt.Templates.Functions()) + for _, variable := range in.pkg.Variables { + lt := newChecked(in.pkg, variable) + index, err := buildTreeIndex(lt, in.dir) if err != nil { return nil, err } diff --git a/internal/mutation/input.go b/internal/mutation/input.go new file mode 100644 index 00000000..ff76da13 --- /dev/null +++ b/internal/mutation/input.go @@ -0,0 +1,74 @@ +package mutation + +import ( + "text/template/parse" + + "github.com/typelate/check" + "golang.org/x/tools/go/packages" + + "github.com/typelate/muxt/internal/load" + "github.com/typelate/muxt/internal/source" +) + +// input is what a run plans from: the package its templates variables are +// declared in, and the directory reported paths are relative to. +// +// loadInput builds one with the go command. Everything after that -- +// traversal, enumeration, the --diff comparison -- reads only this, so a +// test can plan from a package built in memory. +type input struct { + dir string + pkg source.Package +} + +// loadInput loads the package in dir, with its test files when the +// configuration includes test callers. env is the environment the go +// command runs in, nil for the process's own. +func loadInput(dir string, config Configuration, env []string) (input, error) { + var ( + pl []*packages.Package + err error + ) + if config.IncludeTests { + pl, err = load.PackagesWithTests(dir, env) + } else { + _, pl, err = load.PackagesWithEnv(dir, env) + } + if err != nil { + return input{}, err + } + return inputFrom(dir, pl, config.TemplatesVariables) +} + +// inputFrom reads the templates variables from the package in dir among +// pl. +func inputFrom(dir string, pl []*packages.Package, variables []string) (input, error) { + pkg, err := load.Package(dir, pl, variables) + if err != nil { + return input{}, err + } + return input{dir: dir, pkg: pkg}, nil +} + +// checked is one templates variable with the checker built for it, which +// traversal and the validity check share. +type checked struct { + source.Variable + pkg source.Package + global *check.Global +} + +func newChecked(pkg source.Package, variable source.Variable) *checked { + trees := check.FindTreeFunc(func(name string) (*parse.Tree, bool) { + t := variable.Set.Lookup(name) + if t == nil || t.Tree == nil { + return nil, false + } + return t.Tree, true + }) + return &checked{ + Variable: variable, + pkg: pkg, + global: check.NewGlobal(pkg.Types, pkg.Fset, trees, check.Functions(variable.Functions)), + } +} diff --git a/internal/mutation/plan.go b/internal/mutation/plan.go index ee05ffc6..d249ed1b 100644 --- a/internal/mutation/plan.go +++ b/internal/mutation/plan.go @@ -7,13 +7,12 @@ import ( "path/filepath" "slices" "strconv" - "strings" "text/template/parse" "github.com/typelate/check" - "golang.org/x/tools/go/packages" "github.com/typelate/muxt/internal/asteval" + "github.com/typelate/muxt/internal/source" ) // plan is everything decided before a single test is run: which templates @@ -145,65 +144,73 @@ func (s selector) choose(scopes []scope, trims []trim) selection { return chosen } -// newPlan loads the project, walks the templates each ExecuteTemplate -// call reaches, and enumerates the variations available in each. +// newPlan loads the project, then plans from it. func newPlan(config Configuration, workingDirectory string) (*plan, error) { - pl, err := loadPackages(workingDirectory, config.IncludeTests, config.env) + in, err := loadInput(workingDirectory, config, config.env) if err != nil { return nil, err } - - include := func(string) bool { return true } - if config.TemplatePattern != nil { - include = config.TemplatePattern.MatchString - } - - p := &plan{ - seed: config.Seed, - draw: newValues(config.Seed), - maxCases: config.MaxCases, - } - if p.maxCases <= 0 { - p.maxCases = DefaultMaxCases - } // before is what the templates looked like at the --diff revision. // Nil means there is nothing to compare with, so every template // counts as changed. - var before revision + var ( + before revision + diffError string + ) if config.Diff != "" { - p.diff = config.Diff dir, cleanup, err := checkout(workingDirectory, config.Diff) if err != nil { return nil, err } defer cleanup() if before, err = templatesAt(config, dir); err != nil { - p.diffError = err.Error() + diffError = err.Error() } } + return planFrom(config, in, before, diffError) +} + +// planFrom walks the templates each ExecuteTemplate call in the input +// reaches and enumerates the variations available in each. +// +// It is everything a plan decides, from inputs that hold no loader: before +// is what the templates read like at the --diff revision, nil when there +// is none, and diffError why that revision could not be read. +func planFrom(config Configuration, in input, before revision, diffError string) (*plan, error) { + include := func(string) bool { return true } + if config.TemplatePattern != nil { + include = config.TemplatePattern.MatchString + } + + p := &plan{ + seed: config.Seed, + draw: newValues(config.Seed), + maxCases: config.MaxCases, + diff: config.Diff, + diffError: diffError, + } + if p.maxCases <= 0 { + p.maxCases = DefaultMaxCases + } sel := selector{ before: before, seen: make(map[string]struct{}), unchanged: make(map[string]struct{}), include: include, - wd: workingDirectory, + wd: in.dir, } - for _, templatesVariable := range config.TemplatesVariables { - lt, err := asteval.LoadTemplates(workingDirectory, templatesVariable, pl) - if err != nil { - return nil, err - } - functions := lt.Templates.Functions() + for _, variable := range in.pkg.Variables { + lt := newChecked(in.pkg, variable) - index, err := buildTreeIndex(lt, workingDirectory, pl, functions) + index, err := buildTreeIndex(lt, in.dir) if err != nil { return nil, err } chosen := sel.choose(traverse(lt, index)) for _, sc := range chosen.mutate { - p.add(lt, sc, functions, workingDirectory) + p.add(lt, sc, lt.Functions, in.dir) } p.unchanged = append(p.unchanged, chosen.unchanged...) p.trimmed = append(p.trimmed, chosen.trimmed...) @@ -224,7 +231,7 @@ func newPlan(config Configuration, workingDirectory string) (*plan, error) { // add enumerates one template's mutants and files them under the call // that reaches it. -func (p *plan) add(lt *asteval.LoadedTemplates, sc scope, functions check.Functions, workingDirectory string) { +func (p *plan) add(lt *checked, sc scope, functions check.Functions, workingDirectory string) { found, notes := mutantsInScope(sc, functions, p.draw, p.maxCases) report := TemplateReport{ @@ -309,7 +316,7 @@ func (p *plan) add(lt *asteval.LoadedTemplates, sc scope, functions check.Functi // the tests fail with a render error. That failure would be recorded as // the mutation being caught, which is a lie: nothing asserted on the // behaviour, the template just stopped working. -func invalid(lt *asteval.LoadedTemplates, sc scope, mutant Mutant, functions check.Functions) (string, bool) { +func invalid(lt *checked, sc scope, mutant Mutant, functions check.Functions) (string, bool) { mutated := sc.src.mutatedText(mutant.edits) trees, err := asteval.ParseTrees(sc.src.rootName, mutated, sc.src.leftDelim, sc.src.rightDelim, functions) if err != nil { @@ -331,21 +338,22 @@ func invalid(lt *asteval.LoadedTemplates, sc scope, mutant Mutant, functions che // The trees are parsed here rather than taken from the template set so // that every node position is an offset into text this package holds, // which is what a mutation is spliced into. -func buildTreeIndex(lt *asteval.LoadedTemplates, workingDirectory string, pl []*packages.Package, functions check.Functions) (map[string]treeLocation, error) { +func buildTreeIndex(lt *checked, workingDirectory string) (map[string]treeLocation, error) { + functions := lt.Functions // The definitions are gathered before the collector is built: a // source scans its actions as it is constructed, and it can only do // that once the delimiters its file was written with are known, // which is something the definitions say. - var defs []check.Definition - for _, t := range lt.HTML.Templates() { - definition, ok := lt.Templates.FindDefinition(t.Name()) + var defs []source.Definition + for _, t := range lt.Set.Templates() { + definition, ok := lt.Definitions[t.Name()] if !ok { continue } defs = append(defs, definition) } - collector := newSourceCollector(workingDirectory, pl, defs) + collector := newSourceCollector(workingDirectory, defs) for _, definition := range defs { if _, err := collector.add(definition); err != nil { return nil, err @@ -421,46 +429,6 @@ func countActions(node parse.Node) int { } } -// loadPackages loads the working directory's package, optionally -// including its test files so that ExecuteTemplate calls written in tests -// are visible. env is the environment the go command runs in, nil for the -// process's own. -func loadPackages(workingDirectory string, includeTests bool, env []string) ([]*packages.Package, error) { - if !includeTests { - _, pl, err := asteval.LoadPackagesWithEnv(workingDirectory, env) - return pl, err - } - fileSet := token.NewFileSet() - pl, err := packages.Load(&packages.Config{ - Fset: fileSet, - Tests: true, - Mode: packages.NeedModule | packages.NeedTypesInfo | packages.NeedName | - packages.NeedFiles | packages.NeedTypes | packages.NeedSyntax | - packages.NeedEmbedPatterns | packages.NeedEmbedFiles | packages.NeedImports, - Dir: workingDirectory, - Env: env, - }, workingDirectory, "encoding", "fmt", "net/http") - if err != nil { - return nil, err - } - // With Tests set, go list reports the package twice: once as it is - // written and once compiled with its in-package test files. The - // second is the one holding a test's ExecuteTemplate calls, so it - // has to come first when a package is picked by directory. - slices := make([]*packages.Package, 0, len(pl)) - for _, pkg := range pl { - if strings.HasSuffix(pkg.ID, ".test]") { - slices = append(slices, pkg) - } - } - for _, pkg := range pl { - if !strings.HasSuffix(pkg.ID, ".test]") { - slices = append(slices, pkg) - } - } - return slices, nil -} - func relativePosition(workingDirectory string, position token.Position) string { if !position.IsValid() { return "?" diff --git a/internal/mutation/snapshot_test.go b/internal/mutation/snapshot_test.go new file mode 100644 index 00000000..21a46af6 --- /dev/null +++ b/internal/mutation/snapshot_test.go @@ -0,0 +1,154 @@ +package mutation + +import ( + "errors" + "flag" + "os" + "path/filepath" + "slices" + "strings" + "testing" + + "golang.org/x/tools/txtar" + + "github.com/typelate/muxt/internal/load/loadtest" + "github.com/typelate/muxt/internal/muxt" +) + +var update = flag.Bool("update", false, "rewrite the want/ files of the snapshot archives in testdata") + +// TestSnapshots plans a dry run for each archive in testdata and compares +// the report with the archive's want/ files. +// +// An archive holds a case's inputs and what the plan reports, and +// snapshots, in snapshots_test.go, the configuration it runs with: +// +// - Go files and template files are written to a directory and loaded as +// example.com/server by internal/load/loadtest: type checked against +// the stub standard library, without the go command. +// - Files under before/ are the package at the --diff revision, loaded +// the same way, when the configuration names one. +// - want/report.txt is the dry run's report and want/error.txt the error +// planning returned. +// +// Run with -update to rewrite the want/ files, then read the diff. Whether +// the tests catch a mutant is decided by running them, which is what +// run_test.go and the integration suite do. +func TestSnapshots(t *testing.T) { + archives, err := filepath.Glob(filepath.Join("testdata", "*.txtar")) + if err != nil { + t.Fatal(err) + } + for _, archivePath := range archives { + name := strings.TrimSuffix(filepath.Base(archivePath), ".txtar") + if !slices.ContainsFunc(snapshots, func(c snapshotCase) bool { return c.archive == name }) { + t.Errorf("testdata/%s.txtar has no configuration in snapshots", name) + } + } + for _, tt := range snapshots { + t.Run(tt.archive, func(t *testing.T) { + if !tt.config.DryRun { + t.Fatal("a snapshot plans a dry run; the configuration must set DryRun") + } + archivePath := filepath.Join("testdata", tt.archive+".txtar") + archive, err := txtar.ParseFile(archivePath) + if err != nil { + t.Fatal(err) + } + got := dryRunSnapshot(t, tt.config, archive) + if *update { + files := slices.DeleteFunc(slices.Clone(archive.Files), func(file txtar.File) bool { + return strings.HasPrefix(file.Name, "want/") + }) + for _, name := range []string{"error.txt", "report.txt"} { + if text, ok := got[name]; ok { + files = append(files, txtar.File{Name: "want/" + name, Data: []byte(text)}) + } + } + archive.Files = files + if err := os.WriteFile(archivePath, txtar.Format(archive), 0o644); err != nil { + t.Fatal(err) + } + return + } + want := make(map[string]string) + for _, file := range archive.Files { + if name, ok := strings.CutPrefix(file.Name, "want/"); ok { + want[name] = string(file.Data) + } + } + for _, name := range []string{"error.txt", "report.txt"} { + if got[name] != want[name] { + t.Errorf("want/%s differs (run go test -run TestSnapshots -update to rewrite):\n--- got\n%s\n--- want\n%s", name, got[name], want[name]) + } + } + }) + } +} + +// dryRunSnapshot plans and dry runs an archive's package under config. +func dryRunSnapshot(t *testing.T, config Configuration, archive *txtar.Archive) map[string]string { + t.Helper() + current, before := make(map[string]string), make(map[string]string) + for _, file := range archive.Files { + switch { + case strings.HasPrefix(file.Name, "want/"): + case strings.HasPrefix(file.Name, "before/"): + before[strings.TrimPrefix(file.Name, "before/")] = string(file.Data) + default: + current[file.Name] = string(file.Data) + } + } + + got := make(map[string]string) + fail := func(err error) map[string]string { + text := err.Error() + if multiLine, ok := errors.AsType[muxt.MultiLineError](err); ok { + text = multiLine.MultiLineError() + } + got["error.txt"] = text + "\n" + return got + } + + dir := t.TempDir() + in, err := inputFrom(dir, loadtest.Package(t, dir, "example.com/server", current), config.TemplatesVariables) + if err != nil { + return fail(err) + } + + var ( + previous revision + diffError string + ) + if config.Diff != "" { + beforeDir := t.TempDir() + beforeInput, err := inputFrom(beforeDir, loadtest.Package(t, beforeDir, "example.com/server", before), config.TemplatesVariables) + if err == nil { + previous, err = revisionOf(beforeInput) + } + if err != nil { + diffError = err.Error() + } + } + + p, err := planFrom(config, in, previous, diffError) + if err != nil { + return fail(err) + } + report, err := runPlan(p, config, nil, func([]string) (string, error) { + t.Fatal("a dry run ran the baseline") + return "", nil + }, func(string) (Status, error) { + t.Fatal("a dry run ran a mutant") + return "", nil + }) + if err != nil { + return fail(err) + } + var out strings.Builder + if _, err := report.WriteTo(&out); err != nil { + t.Fatal(err) + } + got["report.txt"] = out.String() + return got +} diff --git a/internal/mutation/snapshots_test.go b/internal/mutation/snapshots_test.go new file mode 100644 index 00000000..ba7df317 --- /dev/null +++ b/internal/mutation/snapshots_test.go @@ -0,0 +1,68 @@ +package mutation + +import "regexp" + +// snapshots names the configuration each archive in testdata is planned +// with: the configuration the command line in the comment parses into. +// +// TestCommandLineConfigurations in internal/cli states what command lines +// parse into, with literals like these; search for a literal to find its +// twin. Every configuration here is one a command line can produce: +// rejecting one that cannot work is the command line's job. +type snapshotCase struct { + archive string + config Configuration +} + +var snapshots = []snapshotCase{ + { + // muxt test-template-mutations --dry-run --seed=1 -v + archive: "literal_template", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 -v + archive: "partials_and_trims", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 -v --template-pattern=^footer$ + archive: "template_pattern", + config: Configuration{TemplatesVariables: []string{"templates"}, TemplatePattern: regexp.MustCompile("^footer$"), Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 -v + archive: "skipped_mutants", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 -v --max-cases=2 + archive: "operand_budget", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: 2, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 -v + archive: "delimiters", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 -v --diff=main + archive: "diff", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1, Diff: "main"}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 + archive: "err_no_call_sites", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 + archive: "err_no_mutations", + config: Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, + { + // muxt test-template-mutations --dry-run --seed=1 --use-templates-variable=pages + archive: "err_missing_variable", + config: Configuration{TemplatesVariables: []string{"pages"}, Packages: []string{}, DryRun: true, Seed: 1, SeedSet: true, MaxCases: DefaultMaxCases, Workers: 1}, + }, +} diff --git a/internal/mutation/source.go b/internal/mutation/source.go index a6513210..150f29a1 100644 --- a/internal/mutation/source.go +++ b/internal/mutation/source.go @@ -3,14 +3,14 @@ package mutation import ( "fmt" "go/ast" + "go/parser" "go/token" "path/filepath" "strconv" "strings" "unicode/utf8" - "github.com/typelate/check" - "golang.org/x/tools/go/packages" + "github.com/typelate/muxt/internal/source" ) // templateSource is a file whose bytes hold template text, together with @@ -81,7 +81,7 @@ const spaceChars = " \t\r\n" // an optional trim marker, the word end, another optional marker, and // the right delimiter. Nothing else in it varies, so whatever surrounds // the word is the pair. -func delimiters(text string, definition check.Definition) (left, right string, ok bool) { +func delimiters(text string, definition source.Definition) (left, right string, ok bool) { if !definition.TemplateName.IsValid() { // A template with no define clause has no end clause either. return "", "", false @@ -254,39 +254,37 @@ func literalOffsets(literal, value string) ([]int, error) { return offsets, nil } -// findStringLiteral returns the Go string literal covering offset in the -// named file, which is the literal a template written in Go source was -// written as. -func findStringLiteral(pl []*packages.Package, filename string, offset int) (start, end int, ok bool) { - seen := make(map[*ast.File]struct{}) - for _, pkg := range pl { - for _, file := range pkg.Syntax { - if _, done := seen[file]; done { - continue - } - seen[file] = struct{}{} - tokenFile := pkg.Fset.File(file.Pos()) - if tokenFile == nil || tokenFile.Name() != filename { - continue - } - ast.Inspect(file, func(node ast.Node) bool { - lit, isLit := node.(*ast.BasicLit) - if !isLit || lit.Kind != token.STRING { - return true - } - litStart, litEnd := tokenFile.Offset(lit.Pos()), tokenFile.Offset(lit.End()) - if offset < litStart || offset >= litEnd { - return true - } - // Nested literals do not occur, so the first match is - // the one wanted. - start, end, ok = litStart, litEnd, true - return false - }) - if ok { - return start, end, true - } - } +// findStringLiteral returns the bounds of the Go string literal covering +// offset in the named file's text, which is the literal a template written +// in Go source was written as. +// +// The file is parsed again rather than taken from the loaded package: the +// text is already in hand, parsing one file is cheap, and it leaves the +// plan needing nothing from the loader. A file that does not fully parse +// still yields the literals the parser read. +func findStringLiteral(filename, text string, offset int) (start, end int, ok bool) { + fset := token.NewFileSet() + file, _ := parser.ParseFile(fset, filename, text, parser.SkipObjectResolution) + if file == nil { + return 0, 0, false } - return 0, 0, false + tokenFile := fset.File(file.FileStart) + ast.Inspect(file, func(node ast.Node) bool { + if ok { + return false + } + lit, isLit := node.(*ast.BasicLit) + if !isLit || lit.Kind != token.STRING { + return true + } + litStart, litEnd := tokenFile.Offset(lit.Pos()), tokenFile.Offset(lit.End()) + if offset < litStart || offset >= litEnd { + return true + } + // Nested literals do not occur, so the first match is the one + // wanted. + start, end, ok = litStart, litEnd, true + return false + }) + return start, end, ok } diff --git a/internal/mutation/testdata/delimiters.txtar b/internal/mutation/testdata/delimiters.txtar new file mode 100644 index 00000000..a4f23d81 --- /dev/null +++ b/internal/mutation/testdata/delimiters.txtar @@ -0,0 +1,23 @@ +A literal parsed with other delimiters is read with them. +-- server.go -- +package server + +import ( + "html/template" + "io" +) + +var templates = template.Must(template.New("page").Delims("[[", "]]").Parse(`[[define "note"]]

    [[.]]

    [[end]]`)) + +func RenderNote(w io.Writer, note string) error { + return templates.ExecuteTemplate(w, "note", note) +} +-- want/report.txt -- +1 mutant across 1 template (complexity 1, seed 1) +dry run: no tests were run + +server.go:11:9 ExecuteTemplate "note" (dot: string) + "note" server.go (complexity 1, dot: string) + PEND 8:98 action-zero + +1 mutant, 1 runnable, 0 skipped diff --git a/internal/mutation/testdata/diff.txtar b/internal/mutation/testdata/diff.txtar new file mode 100644 index 00000000..96bb14fa --- /dev/null +++ b/internal/mutation/testdata/diff.txtar @@ -0,0 +1,65 @@ +With --diff, a template whose text and dot are unchanged since the +revision is left alone and reported; one that reads differently is +mutated. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +type Page struct{ Title, Body string } + +func RenderPage(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} +-- page.gohtml -- +{{define "page"}}{{template "header" .}}{{template "body" .}}{{end}} +{{define "header"}}

    {{.Title}}

    {{end}} +{{define "body"}}
    {{.Body}}{{if .Title}}!{{end}}
    {{end}} +-- before/server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +type Page struct{ Title, Body string } + +func RenderPage(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} +-- before/page.gohtml -- +{{define "page"}}{{template "header" .}}{{template "body" .}}{{end}} +{{define "header"}}

    {{.Title}}

    {{end}} +{{define "body"}}
    {{.Body}}
    {{end}} +-- want/report.txt -- +3 mutants across 1 template (complexity 2, seed 1) +2 templates unchanged since main +dry run: no tests were run + +server.go:17:9 ExecuteTemplate "page" (dot: server.Page) + "body" page.gohtml via {{template}} (complexity 2, dot: server.Page) + PEND 3:24 action-zero + PEND 3:33 if-false + PEND 3:33 if-true + +unchanged since main, not mutated: + "page" page.gohtml (dot: server.Page) + "header" page.gohtml (dot: server.Page) + +3 mutants, 3 runnable, 0 skipped diff --git a/internal/mutation/testdata/err_missing_variable.txtar b/internal/mutation/testdata/err_missing_variable.txtar new file mode 100644 index 00000000..99435e63 --- /dev/null +++ b/internal/mutation/testdata/err_missing_variable.txtar @@ -0,0 +1,22 @@ +A templates variable the package does not declare. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +func RenderPage(w io.Writer) error { + return templates.ExecuteTemplate(w, "page", nil) +} +-- page.gohtml -- +{{define "page"}}{{.}}{{end}} +-- want/error.txt -- +variable pages not found in package example.com/server diff --git a/internal/mutation/testdata/err_no_call_sites.txtar b/internal/mutation/testdata/err_no_call_sites.txtar new file mode 100644 index 00000000..c385ff9d --- /dev/null +++ b/internal/mutation/testdata/err_no_call_sites.txtar @@ -0,0 +1,22 @@ +A templates variable no ExecuteTemplate call renders has nothing to +mutate from. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +var _ io.Writer + +-- page.gohtml -- +{{define "page"}}{{.}}{{end}} +-- want/error.txt -- +no templates.ExecuteTemplate calls found: mutation testing needs a call site to know the type of dot diff --git a/internal/mutation/testdata/err_no_mutations.txtar b/internal/mutation/testdata/err_no_mutations.txtar new file mode 100644 index 00000000..06a6470b --- /dev/null +++ b/internal/mutation/testdata/err_no_mutations.txtar @@ -0,0 +1,23 @@ +Templates that were reached but hold no action to vary are an error, not +a pass. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +func RenderPage(w io.Writer) error { + return templates.ExecuteTemplate(w, "page", nil) +} +-- page.gohtml -- +{{define "page"}}

    static

    {{end}} +-- want/error.txt -- +no mutations available: the 1 template(s) reached hold no dynamic or control flow actions, so a run would report every mutant killed without testing anything diff --git a/internal/mutation/testdata/literal_template.txtar b/internal/mutation/testdata/literal_template.txtar new file mode 100644 index 00000000..2668750e --- /dev/null +++ b/internal/mutation/testdata/literal_template.txtar @@ -0,0 +1,31 @@ +A template written as a Go string literal: its actions and branches are +mutated where they are written in the literal. +-- server.go -- +package server + +import ( + "html/template" + "io" +) + +var templates = template.Must(template.New("greeting").Parse(`Hello, {{.Name}}!{{if .Loud}} !!!{{end}}`)) + +type Greeting struct { + Name string + Loud bool +} + +func Render(w io.Writer, greeting Greeting) error { + return templates.ExecuteTemplate(w, "greeting", greeting) +} +-- want/report.txt -- +3 mutants across 1 template (complexity 2, seed 1) +dry run: no tests were run + +server.go:16:9 ExecuteTemplate "greeting" (dot: server.Greeting) + "greeting" server.go (complexity 2, dot: server.Greeting) + PEND 8:70 action-zero + PEND 8:80 if-false + PEND 8:80 if-true + +3 mutants, 3 runnable, 0 skipped diff --git a/internal/mutation/testdata/operand_budget.txtar b/internal/mutation/testdata/operand_budget.txtar new file mode 100644 index 00000000..e3f2c04b --- /dev/null +++ b/internal/mutation/testdata/operand_budget.txtar @@ -0,0 +1,33 @@ +An action with more operand combinations than --max-cases allows +contributes none, and says so where its mutants would have been. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +type Page struct{ A, B, C int } + +func RenderPage(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} +-- page.gohtml -- +{{define "page"}}{{printf "%d %d %d" .A .B .C}}{{end}} +-- want/report.txt -- +2 mutants across 1 template (complexity 1, seed 1) +dry run: no tests were run + +server.go:17:9 ExecuteTemplate "page" (dot: server.Page) + "page" page.gohtml (complexity 1, dot: server.Page) + PEND 1:18 action-zero + SKIP 1:18 operands (3 operands need 7 cases, over --max-cases=2) + +2 mutants, 1 runnable, 1 skipped diff --git a/internal/mutation/testdata/partials_and_trims.txtar b/internal/mutation/testdata/partials_and_trims.txtar new file mode 100644 index 00000000..09630d29 --- /dev/null +++ b/internal/mutation/testdata/partials_and_trims.txtar @@ -0,0 +1,57 @@ +A page renders partials through {{template}}, narrowing dot through a +range. A partial reached twice with the same dot is mutated once, and the +repeat is reported as trimmed. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +type Page struct { + Title string + Items []Item +} + +type Item struct { + Name string + Price int +} + +func RenderPage(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} + +func RenderAgain(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} +-- page.gohtml -- +{{define "page"}}

    {{.Title}}

    {{range .Items}}{{template "item" .}}{{end}}{{template "item" (index .Items 0)}}{{end}} +{{define "item"}}
  • {{.Name}}{{if gt .Price 10}} (sale){{end}}
  • {{end}} +-- want/report.txt -- +7 mutants across 2 templates (complexity 4, seed 1) +dry run: no tests were run + +server.go:25:9 ExecuteTemplate "page" (dot: server.Page) + "page" page.gohtml (complexity 2, dot: server.Page) + PEND 1:22 action-zero + PEND 1:37 range-never + PEND 1:53 template-drop + PEND 1:81 template-drop + "item" page.gohtml via {{template}} (complexity 2, dot: server.Item) + PEND 2:22 action-zero + PEND 2:31 if-false + PEND 2:31 if-true + +trimmed 2 subtrees already mutated with the same dot: + "item" at server.go:25:9, first reached from server.go:25:9 + "page" at server.go:29:9, first reached from server.go:25:9 + +7 mutants, 7 runnable, 0 skipped diff --git a/internal/mutation/testdata/skipped_mutants.txtar b/internal/mutation/testdata/skipped_mutants.txtar new file mode 100644 index 00000000..dd440f91 --- /dev/null +++ b/internal/mutation/testdata/skipped_mutants.txtar @@ -0,0 +1,45 @@ +A mutant that would not type check is skipped rather than run: emptying +the action that declares $inner leaves its use undefined. So is a +condition simplification proves cannot change the decision. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +type Inner struct{ Name string } + +type Page struct { + Inner Inner + Banned bool + Loud bool +} + +func RenderPage(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} +-- page.gohtml -- +{{define "page"}}{{$inner := .Inner}}{{$inner.Name}}{{if or .Banned (and .Banned .Loud)}}B{{end}}{{end}} +-- want/report.txt -- +7 mutants across 1 template (complexity 3, seed 1) +dry run: no tests were run + +server.go:23:9 ExecuteTemplate "page" (dot: server.Page) + "page" page.gohtml (complexity 3, dot: server.Page) + SKIP 1:18 action-empty (does not type check against server.Page) + PEND 1:38 action-empty + PEND 1:53 condition .Banned=false + PEND 1:53 condition .Banned=true + SKIP 1:53 condition-dead (.Loud cannot change the decision) + PEND 1:53 if-false + PEND 1:53 if-true + +7 mutants, 5 runnable, 2 skipped diff --git a/internal/mutation/testdata/template_pattern.txtar b/internal/mutation/testdata/template_pattern.txtar new file mode 100644 index 00000000..0791a054 --- /dev/null +++ b/internal/mutation/testdata/template_pattern.txtar @@ -0,0 +1,33 @@ +--template-pattern limits mutation to the templates it matches; the +traversal still passes through the others to reach them. +-- server.go -- +package server + +import ( + "embed" + "html/template" + "io" +) + +//go:embed *.gohtml +var templateFiles embed.FS + +var templates = template.Must(template.ParseFS(templateFiles, "*.gohtml")) + +type Page struct{ Title string } + +func RenderPage(w io.Writer, page Page) error { + return templates.ExecuteTemplate(w, "page", page) +} +-- page.gohtml -- +{{define "page"}}

    {{.Title}}

    {{template "footer" .}}{{end}} +{{define "footer"}}{{end}} +-- want/report.txt -- +1 mutant across 1 template (complexity 1, seed 1) +dry run: no tests were run + +server.go:17:9 ExecuteTemplate "page" (dot: server.Page) + "footer" page.gohtml via {{template}} (complexity 1, dot: server.Page) + PEND 2:28 action-zero + +1 mutant, 1 runnable, 0 skipped diff --git a/internal/mutation/traverse.go b/internal/mutation/traverse.go index 26d02ff0..2a1d2e8b 100644 --- a/internal/mutation/traverse.go +++ b/internal/mutation/traverse.go @@ -6,8 +6,6 @@ import ( "text/template/parse" "github.com/typelate/check" - - "github.com/typelate/muxt/internal/asteval" ) // callSite is one templates.ExecuteTemplate call, which is where a @@ -80,15 +78,15 @@ type trim struct { // input, or two {{template}} invocations passing the same type, would // produce the same mutants and the same verdicts, so the second is // trimmed. -func traverse(lt *asteval.LoadedTemplates, index map[string]treeLocation) ([]scope, []trim) { +func traverse(lt *checked, index map[string]treeLocation) ([]scope, []trim) { t := &traversal{lt: lt, index: index, visited: make(map[string]callSite)} - for call := range lt.Templates.ExecuteTemplateCalls() { + for _, call := range lt.Calls { site := callSite{ - Position: lt.Package.Fset.Position(call.Call.Pos()), - Template: call.TemplateName, - DataType: call.DataType, + Position: call.Position, + Template: call.Template, + DataType: call.Data, } - t.visit(site, call.TemplateName, call.DataType, false) + t.visit(site, call.Template, call.Data, false) } return t.scopes, t.trimmed } @@ -96,7 +94,7 @@ func traverse(lt *asteval.LoadedTemplates, index map[string]treeLocation) ([]sco // traversal is the state of one walk: what it reads from, and what it has // found so far. type traversal struct { - lt *asteval.LoadedTemplates + lt *checked index map[string]treeLocation // visited records the call site each template and dot was first @@ -148,7 +146,7 @@ type templateCall struct { // // The types come from the checker, which is the only thing that knows how // dot narrows through a range or a with on the way to the invocation. -func templateCalls(lt *asteval.LoadedTemplates, tree *parse.Tree, dot types.Type) []templateCall { +func templateCalls(lt *checked, tree *parse.Tree, dot types.Type) []templateCall { var found []templateCall // A template that does not check still yields the invocations found // before the failure, which is better than none: muxt check is where @@ -161,18 +159,18 @@ func templateCalls(lt *asteval.LoadedTemplates, tree *parse.Tree, dot types.Type // checks reports whether tree type checks with dot, which is how a mutant // is told from a mutation that merely breaks the template. -func checks(lt *asteval.LoadedTemplates, tree *parse.Tree, dot types.Type) bool { +func checks(lt *checked, tree *parse.Tree, dot types.Type) bool { return executeWith(lt, tree, dot, nil) == nil } // executeWith type checks tree with dot, calling inspect on each // {{template}} node it passes, then puts back whatever inspector the // shared checker held before. -func executeWith(lt *asteval.LoadedTemplates, tree *parse.Tree, dot types.Type, inspect func(*parse.TemplateNode, *parse.Tree, types.Type, check.Definition)) error { - saved := lt.Global.InspectTemplateNode - lt.Global.InspectTemplateNode = inspect - defer func() { lt.Global.InspectTemplateNode = saved }() - return check.Execute(lt.Global, tree, dot) +func executeWith(lt *checked, tree *parse.Tree, dot types.Type, inspect func(*parse.TemplateNode, *parse.Tree, types.Type, check.Definition)) error { + saved := lt.global.InspectTemplateNode + lt.global.InspectTemplateNode = inspect + defer func() { lt.global.InspectTemplateNode = saved }() + return check.Execute(lt.global, tree, dot) } // typeKey identifies a type exactly, for deciding whether a template has diff --git a/internal/muxt/call.go b/internal/muxt/call.go index 8853839e..eedcc4f0 100644 --- a/internal/muxt/call.go +++ b/internal/muxt/call.go @@ -10,9 +10,8 @@ import ( "slices" "strings" - "golang.org/x/tools/go/packages" - "github.com/typelate/muxt/internal/astgen" + "github.com/typelate/muxt/internal/source" ) type Argument struct { @@ -119,17 +118,18 @@ const ( ResultShapeError ) -func ResolveCall(def *Definition, templatesPackage *types.Package, receiver *types.Named, pl []*packages.Package) error { +func ResolveCall(def *Definition, pkg source.Package, receiver *types.Named) error { if def.call == nil || def.fun == nil { return nil } - sig, isMethod, args, err := resolveCall(def, def.call, templatesPackage, receiver, pl) + sig, isMethod, args, err := resolveCall(def, def.call, pkg, receiver) if err != nil { return def.finishNameError(err, def.handlerSpan()) } def.sig = sig def.isMethod = isMethod def.Arguments = args + recordPathValueTypes(def.pathValueTypes, args, make(map[string]bool)) shape, err := classifyResultShape(def, typeQualifier(receiver.Obj().Pkg())) if err != nil { // Result-shape errors are about the method contract, so the @@ -140,6 +140,31 @@ func ResolveCall(def *Definition, templatesPackage *types.Package, receiver *typ return def.finishNameError(resolveCallbackShapes(def), def.handlerSpan()) } +// recordPathValueTypes records the type each path parameter parses into. +// +// A parameter is parsed once per request, where the call first passes it +// -- depth first, in argument order -- so that occurrence decides its +// type: the parameter's type, unless a string is assignable to it, in +// which case the value is passed along unparsed and stays a string. +// Later occurrences reuse the first one's value. An sse-prefixed name is a +// render callback wherever it appears, never a parsed value. +func recordPathValueTypes(into map[string]types.Type, args []Argument, seen map[string]bool) { + for _, arg := range args { + switch arg.Type { + case ArgumentTypeCall: + recordPathValueTypes(into, arg.args, seen) + case ArgumentTypeRequestPathValue: + if seen[arg.Identifier] || IsSSEArgument(arg.Identifier) { + continue + } + seen[arg.Identifier] = true + if !isStringAssignable(arg.ParamType) { + into[arg.Identifier] = arg.ParamType + } + } + } +} + // resolveCallbackShapes validates each render-callback argument against the // callback contract — func() error (T = struct{}) or func(T) error — and // records T and whether the callback takes the data argument. On sse routes @@ -274,11 +299,11 @@ func checkNestedCallResultShape(name string, sig *types.Signature, qual types.Qu // definedHere returns a "file:line:col: name is defined here" note for // object, or "" when its source position is unknown (synthesized // methods, for instance, have no position). -func definedHere(pl []*packages.Package, object types.Object) string { - if object == nil || !object.Pos().IsValid() || len(pl) == 0 || pl[0].Fset == nil { +func definedHere(pkg source.Package, object types.Object) string { + if object == nil || !object.Pos().IsValid() || pkg.Fset == nil { return "" } - position := pl[0].Fset.Position(object.Pos()) + position := pkg.Fset.Position(object.Pos()) if position.Filename == "" { return "" } @@ -293,7 +318,7 @@ func definedHere(pl []*packages.Package, object types.Object) string { // When the call identifier is neither a receiver method nor a package-scope // function, its signature is synthesized from the call scope and attached to // the receiver so it appears in the generated RoutesReceiver interface. -func resolveCall(def *Definition, call *ast.CallExpr, templatesPackage *types.Package, receiver *types.Named, pl []*packages.Package) (*types.Signature, bool, []Argument, error) { +func resolveCall(def *Definition, call *ast.CallExpr, pkg source.Package, receiver *types.Named) (*types.Signature, bool, []Argument, error) { fun, ok := call.Fun.(*ast.Ident) if !ok { return nil, false, nil, errAt(call.Fun, "expected a function identifier, got: %s", astgen.Format(call.Fun)) @@ -301,11 +326,11 @@ func resolveCall(def *Definition, call *ast.CallExpr, templatesPackage *types.Pa isMethod := true object, _, _ := types.LookupFieldOrMethod(receiver, true, receiver.Obj().Pkg(), fun.Name) if object == nil { - if m, ok := packageScopeFunc(templatesPackage, fun); ok { + if m, ok := packageScopeFunc(pkg.Types, fun); ok { object = m isMethod = false } else { - ms, err := synthesizeCallSignature(def, call, templatesPackage, receiver, pl) + ms, err := synthesizeCallSignature(def, call, pkg, receiver) if err != nil { return nil, false, nil, err } @@ -316,7 +341,7 @@ func resolveCall(def *Definition, call *ast.CallExpr, templatesPackage *types.Pa } } if call == def.call { - if note := definedHere(pl, object); note != "" { + if note := definedHere(pkg, object); note != "" { def.related = append(def.related, note) } } @@ -354,7 +379,7 @@ func resolveCall(def *Definition, call *ast.CallExpr, templatesPackage *types.Pa args = append(args, Argument{Identifier: argument.Name}) continue } - arg, err := newArgumentFromIdentifier(def, pl, argument, paramType, qual) + arg, err := newArgumentFromIdentifier(def, pkg, argument, paramType, qual) if err != nil { return nil, false, nil, errAtNode(argument, err) } @@ -379,7 +404,7 @@ func resolveCall(def *Definition, call *ast.CallExpr, templatesPackage *types.Pa }) continue } - nestedSig, nestedIsMethod, nestedArgs, err := resolveCall(def, argument, templatesPackage, receiver, pl) + nestedSig, nestedIsMethod, nestedArgs, err := resolveCall(def, argument, pkg, receiver) if err != nil { return nil, false, nil, err } @@ -403,7 +428,7 @@ func resolveCall(def *Definition, call *ast.CallExpr, templatesPackage *types.Pa // defined on the receiver, inferring each parameter type from the argument // scope. Nested calls are resolved (so their own methods are synthesized too) // but do not contribute a parameter, mirroring the pre-hydration generator. -func synthesizeCallSignature(def *Definition, call *ast.CallExpr, templatesPackage *types.Package, receiver *types.Named, pl []*packages.Package) (*types.Signature, error) { +func synthesizeCallSignature(def *Definition, call *ast.CallExpr, pkg source.Package, receiver *types.Named) (*types.Signature, error) { var params []*types.Var hasSSE := false // Each argument becomes a parameter named after it, so a repeated @@ -436,7 +461,7 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, templatesPacka } continue } - tp, ok := DefaultScopeType(pl, def, arg.Name) + tp, ok := DefaultScopeType(pkg, def, arg.Name) if !ok { return nil, errAt(arg, "could not determine a type for %s", arg.Name) } @@ -447,7 +472,7 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, templatesPacka if isCallTo(arg, callWrapperUnmarshalJSON) { // Template-first iteration: without a defined method the decode // target is unknown, so pass the raw payload through. - tp, err := stdlibType(pl, "encoding/json", "RawMessage", false) + tp, err := stdlibType(pkg, "encoding/json", "RawMessage", false) if err != nil { return nil, err } @@ -456,7 +481,7 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, templatesPacka } continue } - if _, _, _, err := resolveCall(def, arg, templatesPackage, receiver, pl); err != nil { + if _, _, _, err := resolveCall(def, arg, pkg, receiver); err != nil { return nil, err } } @@ -468,13 +493,13 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, templatesPacka return types.NewSignatureType(types.NewVar(0, nil, "", receiver.Obj().Type()), nil, nil, types.NewTuple(params...), results, false), nil } -func DefaultScopeType(pl []*packages.Package, def *Definition, argumentIdentifier string) (types.Type, bool) { +func DefaultScopeType(pkg source.Package, def *Definition, argumentIdentifier string) (types.Type, bool) { stdlibType := func(pkgPath, name string, pointer bool) (types.Type, bool) { - pkg, ok := findPackageTypes(pl, pkgPath) + imported, ok := pkg.Import(pkgPath) if !ok { return nil, false } - t := pkg.Scope().Lookup(name).Type() + t := imported.Scope().Lookup(name).Type() if pointer { t = types.NewPointer(t) } @@ -503,34 +528,6 @@ func DefaultScopeType(pl []*packages.Package, def *Definition, argumentIdentifie } } -func findPackageTypes(pl []*packages.Package, pkgPath string) (*types.Package, bool) { - for _, pkg := range pl { - if pkg.Types.Path() == pkgPath { - return pkg.Types, true - } - } - for _, pkg := range pl { - if p, ok := searchImports(pkg.Types, pkgPath); ok { - return p, true - } - } - return nil, false -} - -func searchImports(pt *types.Package, pkgPath string) (*types.Package, bool) { - for _, pkg := range pt.Imports() { - if pkg.Path() == pkgPath { - return pkg, true - } - } - for _, pkg := range pt.Imports() { - if p, ok := searchImports(pkg, pkgPath); ok { - return p, true - } - } - return nil, false -} - // sseCallbackSignature is the func(any) error type synthesized for an sse // argument when the receiver method is not already defined. func sseCallbackSignature() *types.Signature { @@ -580,7 +577,7 @@ func typeQualifier(receiverPkg *types.Package) types.Qualifier { } } -func newArgumentFromIdentifier(def *Definition, pl []*packages.Package, arg *ast.Ident, param types.Type, qual types.Qualifier) (Argument, error) { +func newArgumentFromIdentifier(def *Definition, pkg source.Package, arg *ast.Ident, param types.Type, qual types.Qualifier) (Argument, error) { a := Argument{ Identifier: arg.Name, ParamType: param, @@ -588,36 +585,36 @@ func newArgumentFromIdentifier(def *Definition, pl []*packages.Package, arg *ast switch arg.Name { case TemplateNameScopeIdentifierContext: a.Type = ArgumentTypeRequestContext - if err := isAssignable(pl, param, arg.Name, "context", "Context", false, qual); err != nil { + if err := isAssignable(pkg, param, arg.Name, "context", "Context", false, qual); err != nil { return a, err } case TemplateNameScopeIdentifierForm: a.Type = ArgumentTypeRequestForm - bindings, err := checkFormArgument(def, pl, param, arg.Name, "net/url", "Values", false, qual, false) + bindings, err := checkFormArgument(def, pkg, param, arg.Name, "net/url", "Values", false, qual, false) if err != nil { return a, err } a.formFields = bindings case TemplateNameScopeIdentifierMultipart: a.Type = ArgumentTypeRequestMultipartForm - bindings, err := checkFormArgument(def, pl, param, arg.Name, "mime/multipart", "Form", true, qual, true) + bindings, err := checkFormArgument(def, pkg, param, arg.Name, "mime/multipart", "Form", true, qual, true) if err != nil { return a, err } a.formFields = bindings case TemplateNameScopeIdentifierHTTPRequest: a.Type = ArgumentTypeRequest - if err := isAssignable(pl, param, arg.Name, "net/http", "Request", true, qual); err != nil { + if err := isAssignable(pkg, param, arg.Name, "net/http", "Request", true, qual); err != nil { return a, err } case TemplateNameScopeIdentifierHTTPResponse: a.Type = ArgumentTypeResponse - if err := isAssignable(pl, param, arg.Name, "net/http", "ResponseWriter", false, qual); err != nil { + if err := isAssignable(pkg, param, arg.Name, "net/http", "ResponseWriter", false, qual); err != nil { return a, err } case TemplateNameScopeIdentifierLastEventID: a.Type = ArgumentTypeLastEventID - if err := checkParsedArgument(pl, param, qual); err != nil { + if err := checkParsedArgument(pkg, param, qual); err != nil { return a, err } case TemplateNameScopeIdentifierExecute: @@ -625,13 +622,13 @@ func newArgumentFromIdentifier(def *Definition, pl []*packages.Package, arg *ast a.template = def.template case TemplateNameScopeIdentifierRequestBody: a.Type = ArgumentTypeRequestBody - if err := checkRequestBodyParameter(pl, param, qual); err != nil { + if err := checkRequestBodyParameter(pkg, param, qual); err != nil { return a, err } default: if slices.Contains(def.pathValueNames, arg.Name) { a.Type = ArgumentTypeRequestPathValue - if err := checkParsedArgument(pl, param, qual); err != nil { + if err := checkParsedArgument(pkg, param, qual); err != nil { return a, err } return a, nil @@ -670,20 +667,20 @@ func newArgumentFromIdentifier(def *Definition, pl []*packages.Package, arg *ast return a, nil } -func stdlibType(pl []*packages.Package, pkgPath, name string, pointer bool) (types.Type, error) { - pkg, ok := findPackageTypes(pl, pkgPath) +func stdlibType(pkg source.Package, pkgPath, name string, pointer bool) (types.Type, error) { + imported, ok := pkg.Import(pkgPath) if !ok { return nil, fmt.Errorf("could not find package %q for %s", pkgPath, name) } - t := pkg.Scope().Lookup(name).Type() + t := imported.Scope().Lookup(name).Type() if pointer { t = types.NewPointer(t) } return t, nil } -func isAssignable(pl []*packages.Package, paramType types.Type, argName, packagePath, identifier string, pointer bool, qual types.Qualifier) error { - at, err := stdlibType(pl, packagePath, identifier, pointer) +func isAssignable(pkg source.Package, paramType types.Type, argName, packagePath, identifier string, pointer bool, qual types.Qualifier) error { + at, err := stdlibType(pkg, packagePath, identifier, pointer) if err != nil { return err } @@ -759,8 +756,8 @@ const ( // checkRequestBodyParameter requires the parameter bound to the reserved body // argument to be exactly io.Reader. The request body is a single-use stream, // so the method must not be able to assume more than one read. -func checkRequestBodyParameter(pl []*packages.Package, param types.Type, qual types.Qualifier) error { - readerType, err := stdlibType(pl, "io", "Reader", false) +func checkRequestBodyParameter(pkg source.Package, param types.Type, qual types.Qualifier) error { + readerType, err := stdlibType(pkg, "io", "Reader", false) if err != nil { return err } diff --git a/internal/muxt/call_internal_test.go b/internal/muxt/call_internal_test.go index 7b05d952..5cbb1384 100644 --- a/internal/muxt/call_internal_test.go +++ b/internal/muxt/call_internal_test.go @@ -7,6 +7,7 @@ import ( "testing" "github.com/typelate/muxt/internal/astgen" + "github.com/typelate/muxt/internal/source" ) func mustParseCall(t *testing.T, src string) *ast.CallExpr { @@ -67,7 +68,7 @@ func TestDefinitionsBodyArgumentErrors(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { ts := template.Must(template.New("").Parse(tt.template)) - _, err := Definitions(ts, "templates", nil) + _, err := Definitions(source.Variable{Name: "templates", Set: ts}) if err == nil { t.Fatalf("Definitions(%q) = nil error, want %q", tt.template, tt.wantErr) } @@ -161,7 +162,7 @@ func TestRewriteSignalsArguments(t *testing.T) { func TestDefinitionsSignals(t *testing.T) { t.Run("signals marks the definition", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "POST /search Save(ctx, signals)"}}{{end}}`)) - defs, err := Definitions(ts, "templates", nil) + defs, err := Definitions(source.Variable{Name: "templates", Set: ts}) if err != nil { t.Fatal(err) } @@ -174,7 +175,7 @@ func TestDefinitionsSignals(t *testing.T) { }) t.Run("a signals path wildcard keeps its path-value meaning", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "GET /s/{signals} Show(ctx, signals)"}}{{end}}`)) - defs, err := Definitions(ts, "templates", nil) + defs, err := Definitions(source.Variable{Name: "templates", Set: ts}) if err != nil { t.Fatal(err) } @@ -201,7 +202,7 @@ func TestIsSignalsCallbackArgument(t *testing.T) { func TestDefinitionsSignalsCallback(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "GET /board sse(Stream(ctx, execute, countsSignals))"}}{{end}}`)) - defs, err := Definitions(ts, "templates", nil) + defs, err := Definitions(source.Variable{Name: "templates", Set: ts}) if err != nil { t.Fatal(err) } diff --git a/internal/muxt/call_test.go b/internal/muxt/call_test.go index eccf9b6f..5149c056 100644 --- a/internal/muxt/call_test.go +++ b/internal/muxt/call_test.go @@ -1,28 +1,21 @@ package muxt import ( - "go/token" "go/types" "html/template" + "os" + "path/filepath" "testing" "github.com/stretchr/testify/require" - "golang.org/x/tools/go/packages" + + "github.com/typelate/muxt/internal/source" + "github.com/typelate/muxt/internal/typestest" ) func TestArgument(t *testing.T) { - // The testdata module is never part of a workspace; a GOWORK from the - // invoking environment must not leak into its package loading. - t.Setenv("GOWORK", "off") - fileSet := token.NewFileSet() - packageList, err := packages.Load(&packages.Config{ - Fset: fileSet, - Mode: packages.NeedModule | packages.NeedTypesInfo | packages.NeedName | packages.NeedFiles | packages.NeedTypes | packages.NeedSyntax | packages.NeedEmbedPatterns | packages.NeedEmbedFiles | packages.NeedImports, - Dir: "testdata/example", - }, ".") - require.NoError(t, err) - - examplePkg := packageList[0].Types + examplePkg := exampleTypes(t) + pkg := source.Package{Fset: typestest.FileSet, Types: examplePkg, Imports: typestest.Packages()} require.NotNil(t, examplePkg) httpPkg := findImport(examplePkg, "net/http") @@ -474,13 +467,13 @@ func TestArgument(t *testing.T) { } { t.Run(tc.Name, func(t *testing.T) { ts := template.Must(template.New("").Parse(tc.Template)) - defs, err := Definitions(ts, "templates", nil) + defs, err := Definitions(source.Variable{Name: "templates", Set: ts}) if err != nil { t.Fatal(err) } for i := range defs { - err = ResolveCall(&defs[i], examplePkg, tc.Receiver, packageList) + err = ResolveCall(&defs[i], pkg, tc.Receiver) if err != nil { break } @@ -517,3 +510,18 @@ func findImport(example *types.Package, pkg string) *types.Package { } return nil } + +// exampleTypes type checks the package in testdata/example against the +// stub standard library. +func exampleTypes(t *testing.T) *types.Package { + t.Helper() + files := make(map[string]string) + for _, name := range []string{"functions.go", "methods.go"} { + src, err := os.ReadFile(filepath.Join("testdata", "example", name)) + require.NoError(t, err) + files[name] = string(src) + } + pkg, err := typestest.Check("example.com", files) + require.NoError(t, err) + return pkg +} diff --git a/internal/muxt/definition.go b/internal/muxt/definition.go index faed37ed..eab11d25 100644 --- a/internal/muxt/definition.go +++ b/internal/muxt/definition.go @@ -16,16 +16,15 @@ import ( "strings" "text/template/parse" - "github.com/typelate/check" - "github.com/typelate/muxt/internal/astgen" + "github.com/typelate/muxt/internal/source" ) -// Definitions parses route definitions from the template names in ts. -// The optional definitions finder locates each template's define clause -// so template name errors carry a file position; pass nil when the -// source locations are unknown. -func Definitions(ts *template.Template, templatesVariable string, definitions check.DefinitionFinder) ([]Definition, error) { +// Definitions parses route definitions from the template names in a +// templates variable's set. When the variable locates a template's +// definition, errors about its name carry a file position. +func Definitions(variable source.Variable) ([]Definition, error) { + ts, templatesVariable := variable.Set, variable.Name var defs []Definition type nameFailure struct { def Definition @@ -37,14 +36,8 @@ func Definitions(ts *template.Template, templatesVariable string, definitions ch if !ok { continue } - if definitions != nil { - if d, found := definitions.FindDefinition(t.Name()); found && d.TemplateName.IsValid() { - // The span includes the quotes; the name starts one byte in. - pos := d.TemplateName.Position - pos.Column++ - pos.Offset++ - mt.namePosition = pos - } + if pos, found := variable.NamePosition(t.Name()); found { + mt.namePosition = pos } if err != nil { // Collect every malformed name so one run reports them all. @@ -468,7 +461,9 @@ func (def Definition) SignalsCallback() (string, bool) { return def.signalsCallback, def.signalsCallback != "" } -func (def Definition) SetArgumentType(name string, tp types.Type) { def.pathValueTypes[name] = tp } +// ArgumentType returns the type a path parameter parses into. It is +// unset for a parameter that is passed along as the string it arrived +// as, or is not passed to the call at all. func (def Definition) ArgumentType(name string) (types.Type, bool) { tp, ok := def.pathValueTypes[name] return tp, ok diff --git a/internal/muxt/definition_internal_test.go b/internal/muxt/definition_internal_test.go index 7fddb5e9..8755e29a 100644 --- a/internal/muxt/definition_internal_test.go +++ b/internal/muxt/definition_internal_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/typelate/muxt/internal/source" ) func TestTemplateName_ByPathThenMethod(t *testing.T) { @@ -222,14 +223,14 @@ func TestDefinitionsErrorIncludesSourceFile(t *testing.T) { t.Run("template parsed from a file", func(t *testing.T) { // ParseFS and ParseFiles record the file name as the tree's ParseName. ts := template.Must(template.New("template.gohtml").Parse(`{{define "OPTIONS / F()"}}{{end}}`)) - _, err := Definitions(ts, "templates", nil) + _, err := Definitions(source.Variable{Name: "templates", Set: ts}) require.ErrorContains(t, err, "template.gohtml: OPTIONS method not allowed") }) t.Run("template defined without a file", func(t *testing.T) { // Parse-defined templates carry their own name as ParseName; the // error must not be prefixed with the template name. ts := template.Must(template.New("OPTIONS / F()").Parse(``)) - _, err := Definitions(ts, "templates", nil) + _, err := Definitions(source.Variable{Name: "templates", Set: ts}) require.EqualError(t, err, "OPTIONS method not allowed; allowed methods: GET, POST, PUT, PATCH, and DELETE") }) } diff --git a/internal/muxt/definition_test.go b/internal/muxt/definition_test.go index 2fb4476a..6c0bb45f 100644 --- a/internal/muxt/definition_test.go +++ b/internal/muxt/definition_test.go @@ -10,12 +10,13 @@ import ( "github.com/stretchr/testify/require" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) func TestDefinitions(t *testing.T) { t.Run("when one of the template names is a malformed pattern", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "HEAD /"}}{{end}}`)) - _, err := muxt.Definitions(ts, "ts", nil) + _, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.Error(t, err) }) } @@ -23,7 +24,7 @@ func TestDefinitions(t *testing.T) { func TestCheckPathMethodCollisions(t *testing.T) { t.Run("when two handlers differ only in the case of the first letter", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "GET /items list(ctx)"}}{{end}}{{define "GET /items/{id} List(ctx, id)"}}{{end}}`)) - defs, err := muxt.Definitions(ts, "ts", nil) + defs, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.NoError(t, err) require.ErrorContains(t, muxt.CheckPathMethodCollisions(defs), `TemplateRoutePaths method name collision: handlers "list" and "List" both produce method "List"`) }) @@ -31,13 +32,13 @@ func TestCheckPathMethodCollisions(t *testing.T) { // 一覧 (Japanese "list") has no uppercase form, so no exported // TemplateRoutePaths method name can be derived from it. ts := template.Must(template.New("").Parse(`{{define "GET /items 一覧(ctx)"}}{{end}}`)) - defs, err := muxt.Definitions(ts, "ts", nil) + defs, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.NoError(t, err) require.ErrorContains(t, muxt.CheckPathMethodCollisions(defs), `cannot export identifier "一覧" for TemplateRoutePaths method: first character '一' has no uppercase form`) }) t.Run("when handlers produce distinct method names", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "GET /items List(ctx)"}}{{end}}{{define "GET /items/{id} Show(ctx, id)"}}{{end}}`)) - defs, err := muxt.Definitions(ts, "ts", nil) + defs, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.NoError(t, err) require.NoError(t, muxt.CheckPathMethodCollisions(defs)) }) @@ -46,7 +47,7 @@ func TestCheckPathMethodCollisions(t *testing.T) { func TestCheckForDuplicatePatterns(t *testing.T) { t.Run("when the pattern is not unique", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "GET / F1()"}}a{{end}} {{define "GET / F2()"}}b{{end}}`)) - definitions, err := muxt.Definitions(ts, "ts", nil) + definitions, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.NoError(t, err) require.Len(t, definitions, 2) for _, def := range definitions { @@ -57,7 +58,7 @@ func TestCheckForDuplicatePatterns(t *testing.T) { t.Run("ensure hosts are normalized", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define "GET example.com/ F1()"}}a{{end}} {{define "GET Example.COM/ F2()"}}b{{end}}`)) - definitions, err := muxt.Definitions(ts, "ts", nil) + definitions, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.NoError(t, err) require.Len(t, definitions, 2) for _, def := range definitions { @@ -68,7 +69,7 @@ func TestCheckForDuplicatePatterns(t *testing.T) { t.Run("ensure paths are normalized", func(t *testing.T) { ts := template.Must(template.New("").Parse(`{{define " /abc"}}a{{end}} {{define "/abc "}}b{{end}}`)) - definitions, err := muxt.Definitions(ts, "ts", nil) + definitions, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.NoError(t, err) require.Len(t, definitions, 2) for _, def := range definitions { @@ -102,7 +103,7 @@ func TestCheckForDuplicatePatterns(t *testing.T) { "a.gohtml": &fstest.MapFile{Data: []byte(`{{define "GET / F1()"}}a{{end}}`)}, } ts := template.Must(template.ParseFS(fsys, "*.gohtml")) - definitions, err := muxt.Definitions(ts, "ts", nil) + definitions, err := muxt.Definitions(source.Variable{Name: "ts", Set: ts}) require.NoError(t, err) require.Len(t, definitions, 3) for range 8 { diff --git a/internal/muxt/name_error_test.go b/internal/muxt/name_error_test.go index d8c749d7..a5898da7 100644 --- a/internal/muxt/name_error_test.go +++ b/internal/muxt/name_error_test.go @@ -8,38 +8,23 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/typelate/check" + "github.com/typelate/muxt/internal/source" ) -// definitionAt is a check.DefinitionFinder reporting a fixed name literal position. -type definitionAt struct { - name string - pos token.Position -} - -func (d definitionAt) FindDefinition(name string) (check.Definition, bool) { - if name != d.name { - return check.Definition{}, false - } - return check.Definition{ - Name: name, - TemplateName: check.Span{Position: d.pos, Length: len(name) + 2}, - }, true -} - func TestDefinitionsErrorPosition(t *testing.T) { const name = "OPTIONS / F()" ts := template.Must(template.New("index.gohtml").Parse(`{{define "` + name + `"}}{{end}}`)) - finder := definitionAt{name: name, pos: token.Position{ - Filename: "index.gohtml", - Offset: 9, - Line: 1, - Column: 10, + variable := source.Variable{Name: "templates", Set: ts, Definitions: map[string]source.Definition{ + name: { + Name: name, + // The quoted name starts at column 10, so the name itself + // starts one byte in, at column 11. + TemplateName: source.Span{Position: token.Position{Filename: "index.gohtml", Offset: 9, Line: 1, Column: 10}, Length: len(name) + 2}, + }, }} - _, err := Definitions(ts, "templates", finder) - // The name literal's content starts one byte past the opening quote at - // column 11; the failing METHOD segment starts at the first byte of the name. + _, err := Definitions(variable) + // The failing METHOD segment starts at the first byte of the name. require.EqualError(t, err, "index.gohtml:1:11: OPTIONS method not allowed; allowed methods: GET, POST, PUT, PATCH, and DELETE") nameErr, ok := err.(*NameError) diff --git a/internal/muxt/path_value_types_test.go b/internal/muxt/path_value_types_test.go new file mode 100644 index 00000000..47b5502e --- /dev/null +++ b/internal/muxt/path_value_types_test.go @@ -0,0 +1,76 @@ +package muxt_test + +import ( + "go/types" + "html/template" + "testing" + + "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" + "github.com/typelate/muxt/internal/typestest" +) + +// pathValueReceiver declares methods whose parameters a path value is +// passed to, one parameter type each. +const pathValueReceiver = `package server + +import "time" + +type ID string + +type T struct{} + +func (T) Int(int) any { return nil } +func (T) String(string) any { return nil } +func (T) Any(any) any { return nil } +func (T) Time(time.Time) any { return nil } +func (T) IntString(int, string) any { return nil } +func (T) StringInt(string, int) any { return nil } +func (T) Outer(any, int) any { return nil } +func (T) Inner(int) int { return 0 } +func (T) Pair(int, int) any { return nil } +` + +// TestPathValueTypes states which type a path parameter parses into: the +// parameter type of the first place the call passes it, unless a string +// can be passed there as it is. +func TestPathValueTypes(t *testing.T) { + for _, tt := range []struct { + name string + template string + param string + want string // "" when the parameter stays the string it arrived as + }{ + {name: "parsed into an int parameter", template: "GET /{id} Int(id)", param: "id", want: "int"}, + {name: "a string parameter needs no parsing", template: "GET /{id} String(id)", param: "id"}, + {name: "a string is assignable to any", template: "GET /{id} Any(id)", param: "id"}, + {name: "a text unmarshaler", template: "GET /{at} Time(at)", param: "at", want: "time.Time"}, + {name: "the first occurrence decides when it parses", template: "GET /{id} IntString(id, id)", param: "id", want: "int"}, + {name: "the first occurrence decides when it does not", template: "GET /{id} StringInt(id, id)", param: "id"}, + {name: "a nested call is walked where it is passed", template: "GET /{id} Outer(Inner(id), id)", param: "id", want: "int"}, + {name: "two parameters", template: "GET /{a}/{b} Pair(a, b)", param: "b", want: "int"}, + {name: "a parameter the call does not pass", template: "GET /{a}/{b} Int(a)", param: "b"}, + {name: "an sse-prefixed name is a callback, not a value", template: "GET /{sseID} Int(sseID)", param: "sseID"}, + } { + t.Run(tt.name, func(t *testing.T) { + pkg := typestest.MustCheck(t, "example.com/server", pathValueReceiver) + receiver := pkg.Scope().Lookup("T").Type().(*types.Named) + ts := template.Must(template.New("").Parse(`{{define "` + tt.template + `"}}{{end}}`)) + defs, err := muxt.Definitions(source.Variable{Name: "templates", Set: ts}) + if err != nil { + t.Fatal(err) + } + if err := muxt.ResolveCall(&defs[0], source.Package{Fset: typestest.FileSet, Types: pkg, Imports: typestest.Packages()}, receiver); err != nil { + t.Fatal(err) + } + tp, ok := defs[0].ArgumentType(tt.param) + got := "" + if ok { + got = types.TypeString(tp, (*types.Package).Name) + } + if got != tt.want { + t.Errorf("ArgumentType(%q) = %q, want %q", tt.param, got, tt.want) + } + }) + } +} diff --git a/internal/muxt/testdata/example/go.mod b/internal/muxt/testdata/example/go.mod deleted file mode 100644 index 5c50f802..00000000 --- a/internal/muxt/testdata/example/go.mod +++ /dev/null @@ -1,3 +0,0 @@ -module example.com - -go 1.26 \ No newline at end of file diff --git a/internal/muxt/unmarshal.go b/internal/muxt/unmarshal.go index 7b054cd1..8f350d03 100644 --- a/internal/muxt/unmarshal.go +++ b/internal/muxt/unmarshal.go @@ -8,9 +8,9 @@ import ( "strings" "github.com/typelate/dom" + "github.com/typelate/muxt/internal/source" "golang.org/x/net/html" "golang.org/x/net/html/atom" - "golang.org/x/tools/go/packages" ) // UnmarshalMethod identifies how a request value (path value, lastEventID @@ -39,9 +39,10 @@ const ( // UnmarshalMethodFor classifies how tp parses from its string form: a basic // type parsed with strconv (matched by name, so the byte and rune aliases are // not supported), or a named type whose pointer implements -// encoding.TextUnmarshaler. asteval.LoadPackages always loads the encoding -// package (like fmt), so detection needs nothing from user code. -func UnmarshalMethodFor(pl []*packages.Package, tp types.Type) UnmarshalMethod { +// encoding.TextUnmarshaler. The encoding package is found through pkg.Import; +// load.Packages always loads it (like fmt), so detection needs nothing from +// user code. +func UnmarshalMethodFor(pkg source.Package, tp types.Type) UnmarshalMethod { switch t := tp.(type) { case *types.Basic: switch t.Name() { @@ -75,7 +76,7 @@ func UnmarshalMethodFor(pl []*packages.Package, tp types.Type) UnmarshalMethod { return UnmarshalFloat64 } case *types.Named: - if encPkg, ok := findPackageTypes(pl, "encoding"); ok { + if encPkg, ok := pkg.Import("encoding"); ok { textUnmarshaler := encPkg.Scope().Lookup("TextUnmarshaler").Type().Underlying().(*types.Interface) if types.Implements(types.NewPointer(t), textUnmarshaler) { return UnmarshalTextUnmarshaler @@ -108,21 +109,27 @@ func unsupportedTypeError(tp types.Type, qual types.Qualifier, supported string) // checkUnmarshalable reports whether tp parses from a form field's // string value. -func checkUnmarshalable(pl []*packages.Package, tp types.Type, qual types.Qualifier) error { - if UnmarshalMethodFor(pl, tp) != UnmarshalUnsupported { +func checkUnmarshalable(pkg source.Package, tp types.Type, qual types.Qualifier) error { + if UnmarshalMethodFor(pkg, tp) != UnmarshalUnsupported { return nil } return unsupportedTypeError(tp, qual, supportedUnmarshalFieldTypes) } +// isStringAssignable reports whether a request string can be passed to a +// parameter of type tp without parsing it. +func isStringAssignable(tp types.Type) bool { + return types.AssignableTo(types.Universe.Lookup("string").Type(), tp) +} + // checkParsedArgument validates a path value or lastEventID parameter: // it either receives the raw string or parses from one. Floats are // rejected here even though form fields accept them. -func checkParsedArgument(pl []*packages.Package, paramType types.Type, qual types.Qualifier) error { - if types.AssignableTo(types.Universe.Lookup("string").Type(), paramType) { +func checkParsedArgument(pkg source.Package, paramType types.Type, qual types.Qualifier) error { + if isStringAssignable(paramType) { return nil } - switch UnmarshalMethodFor(pl, paramType) { + switch UnmarshalMethodFor(pkg, paramType) { case UnmarshalUnsupported, UnmarshalFloat32, UnmarshalFloat64: return unsupportedTypeError(paramType, qual, supportedUnmarshalTypes) default: @@ -170,8 +177,8 @@ type FieldBinding struct { // field (nil in raw mode). Struct fields must be a supported scalar or slice // of scalars; multipart structs may also bind *multipart.FileHeader and // []*multipart.FileHeader fields. -func checkFormArgument(def *Definition, pl []*packages.Package, paramType types.Type, argName, packagePath, identifier string, pointer bool, qual types.Qualifier, allowFileFields bool) ([]FieldBinding, error) { - at, err := stdlibType(pl, packagePath, identifier, pointer) +func checkFormArgument(def *Definition, pkg source.Package, paramType types.Type, argName, packagePath, identifier string, pointer bool, qual types.Qualifier, allowFileFields bool) ([]FieldBinding, error) { + at, err := stdlibType(pkg, packagePath, identifier, pointer) if err != nil { return nil, err } @@ -182,13 +189,13 @@ func checkFormArgument(def *Definition, pl []*packages.Package, paramType types. if !ok { return nil, fmt.Errorf("expected %s parameter type to be a struct", argName) } - return formStructBindings(def, pl, st, argName, qual, allowFileFields) + return formStructBindings(def, pkg, st, argName, qual, allowFileFields) } -func formStructBindings(def *Definition, pl []*packages.Package, st *types.Struct, argName string, qual types.Qualifier, allowFileFields bool) ([]FieldBinding, error) { +func formStructBindings(def *Definition, pkg source.Package, st *types.Struct, argName string, qual types.Qualifier, allowFileFields bool) ([]FieldBinding, error) { var fileHeaderPtr types.Type if allowFileFields { - if mp, ok := findPackageTypes(pl, "mime/multipart"); ok { + if mp, ok := pkg.Import("mime/multipart"); ok { if obj := mp.Scope().Lookup("FileHeader"); obj != nil { fileHeaderPtr = types.NewPointer(obj.Type()) } @@ -224,10 +231,10 @@ func formStructBindings(def *Definition, pl []*packages.Package, st *types.Struc return nil, err } fb.Validations = validations - if err := checkUnmarshalable(pl, fb.Elem, qual); err != nil { + if err := checkUnmarshalable(pkg, fb.Elem, qual); err != nil { return nil, fmt.Errorf("failed to generate parse statements for %s field %s: %w", argName, field.Name(), err) } - fb.Method = UnmarshalMethodFor(pl, fb.Elem) + fb.Method = UnmarshalMethodFor(pkg, fb.Elem) bindings = append(bindings, fb) } return bindings, nil diff --git a/internal/source/source.go b/internal/source/source.go new file mode 100644 index 00000000..63f87142 --- /dev/null +++ b/internal/source/source.go @@ -0,0 +1,146 @@ +// Package source holds a Go package as muxt reads it: its types, and each +// templates variable it declares with where its templates were written and +// where they are executed. +// +// internal/load builds a Package from a go/packages load; that is the only +// step that runs the go command. Everything muxt does after it -- route +// resolution, generation, type checking templates, planning mutations -- +// reads a Package. It is plain data, with no functions to call and no +// loader behind it, so a test can write one as a literal or build one from +// source type checked in memory. +package source + +import ( + "go/token" + "go/types" + "html/template" + "text/template/parse" +) + +// Package is a loaded Go package. +type Package struct { + // Fset positions every object in Types and Imports, and every + // position in the variables. + Fset *token.FileSet + + // Types is the package: the one declaring the templates variables, + // whose package-scope functions a template name may call. + Types *types.Package + + // Imports holds, by import path, the packages loaded with Types and + // every package they import: where the standard library types a + // route argument binds to (net/http, context, encoding, ...) are + // found. It may be nil, in which case Import searches Types' own + // imports. + Imports map[string]*types.Package + + // Variables are the templates variables that were asked for, in the + // order they were named. + Variables []Variable +} + +// Import finds the package with path: Types itself, then Imports, then -- +// when Imports is nil -- Types' imports, transitively. +func (pkg Package) Import(path string) (*types.Package, bool) { + if pkg.Types != nil && pkg.Types.Path() == path { + return pkg.Types, true + } + if pkg.Imports != nil { + imported, ok := pkg.Imports[path] + return imported, ok + } + if pkg.Types == nil { + return nil, false + } + return SearchImports(pkg.Types, path) +} + +// SearchImports looks for the package with path among the imports of pt, +// the direct imports before any of theirs. +func SearchImports(pt *types.Package, path string) (*types.Package, bool) { + for _, imported := range pt.Imports() { + if imported.Path() == path { + return imported, true + } + } + for _, imported := range pt.Imports() { + if found, ok := SearchImports(imported, path); ok { + return found, true + } + } + return nil, false +} + +// Variable is one package-level templates variable. +type Variable struct { + // Name is the variable's name. + Name string + + // Set is the template set the variable holds. A text/template + // variable is carried as an html/template value with the same trees: + // muxt reads names and trees and never executes the set. + Set *template.Template + + // Functions are the functions a template in the set may call: the + // builtins and those registered with Funcs. + Functions map[string]*types.Signature + + // Funcs are only the functions registered with Funcs calls in the + // variable's construction chain. + Funcs map[string]*types.Signature + + // Definitions locate where each template in the set was written, by + // template name. A template whose source is unknown has none. + Definitions map[string]Definition + + // Calls are the variable's ExecuteTemplate calls with a string + // literal template name, in file order. + Calls []Call +} + +// NamePosition reports where a template's name begins in the source that +// defines it: the first byte inside the name's quotes. +func (v Variable) NamePosition(templateName string) (token.Position, bool) { + definition, ok := v.Definitions[templateName] + if !ok || !definition.TemplateName.IsValid() { + return token.Position{}, false + } + pos := definition.TemplateName.Position + pos.Column++ + pos.Offset++ + return pos, true +} + +// Definition locates the text defining one template. +type Definition struct { + // Name is the template's name. + Name string + + // Define spans the define or block clause, End the clause that closes + // it, and TemplateName the quoted name in the define clause. For a + // template with no define clause -- the one a parsed text itself + // carries -- TemplateName and End are unset and Define spans the text. + Define, End, TemplateName Span + + // Tree is the template's parse tree. + Tree *parse.Tree +} + +// Span is a run of bytes in a file. +type Span struct { + token.Position + Length int +} + +// Call is one templatesVariable.ExecuteTemplate(w, name, data) call. +type Call struct { + // Position is where the call is written. + Position token.Position + + // Template is the template name the call passes. + Template string + + // Data is the type of the data argument, the type of dot the + // template is executed with. + Data types.Type +} diff --git a/internal/typestest/stubs.go b/internal/typestest/stubs.go new file mode 100644 index 00000000..d2051e29 --- /dev/null +++ b/internal/typestest/stubs.go @@ -0,0 +1,279 @@ +package typestest + +// stubs maps an import path to the source of its stub package. The +// declarations follow the standard library's; bodies return zero values. +var stubs = map[string]string{ + "context": `package context + +import "time" + +type Context interface { + Deadline() (deadline time.Time, ok bool) + Done() <-chan struct{} + Err() error + Value(key any) any +} + +func Background() Context { return nil } +`, + + "encoding": `package encoding + +type TextMarshaler interface { + MarshalText() (text []byte, err error) +} + +type TextUnmarshaler interface { + UnmarshalText(text []byte) error +} +`, + + "encoding/json": `package json + +type RawMessage []byte + +func (m RawMessage) MarshalJSON() ([]byte, error) { return nil, nil } + +func (m *RawMessage) UnmarshalJSON(data []byte) error { return nil } + +func Marshal(v any) ([]byte, error) { return nil, nil } + +func Unmarshal(data []byte, v any) error { return nil } +`, + + "errors": `package errors + +func New(text string) error { return nil } + +func Join(errs ...error) error { return nil } +`, + + "fmt": `package fmt + +type Stringer interface { + String() string +} + +func Sprint(a ...any) string { return "" } + +func Sprintf(format string, a ...any) string { return "" } + +func Sprintln(a ...any) string { return "" } + +func Errorf(format string, a ...any) error { return nil } +`, + + "embed": `package embed + +import "io/fs" + +type FS struct{} + +func (f FS) Open(name string) (fs.File, error) { return nil, nil } +`, + + "html/template": `package template + +import ( + "fmt" + "io" + "io/fs" +) + +type Template struct{} + +type FuncMap map[string]any + +func New(name string) *Template { return nil } + +func Must(t *Template, err error) *Template { return t } + +func ParseFS(fsys fs.FS, patterns ...string) (*Template, error) { return nil, nil } + +func (t *Template) New(name string) *Template { return t } + +func (t *Template) Parse(text string) (*Template, error) { return t, nil } + +func (t *Template) ParseFS(fsys fs.FS, patterns ...string) (*Template, error) { return t, nil } + +func (t *Template) Funcs(funcMap FuncMap) *Template { return t } + +func (t *Template) Delims(left, right string) *Template { return t } + +func (t *Template) Option(opt ...string) *Template { return t } + +func (t *Template) ExecuteTemplate(wr io.Writer, name string, data any) error { return nil } + +func HTMLEscaper(args ...any) string { return "" } + +func JSEscaper(args ...any) string { return "" } + +func URLQueryEscaper(args ...any) string { return "" } + +// The real package reaches fmt through its imports; check.DefaultFunctions +// finds print, printf and println there. +var _ fmt.Stringer +`, + + "io/fs": `package fs + +type FileInfo interface { + Name() string + Size() int64 + IsDir() bool +} + +type File interface { + Stat() (FileInfo, error) + Read([]byte) (int, error) + Close() error +} + +type FS interface { + Open(name string) (File, error) +} +`, + + "io": `package io + +type Reader interface { + Read(p []byte) (n int, err error) +} + +type Writer interface { + Write(p []byte) (n int, err error) +} + +type Closer interface { + Close() error +} + +type ReadCloser interface { + Reader + Closer +} + +func WriteString(w Writer, s string) (n int, err error) { return 0, nil } +`, + + "mime/multipart": `package multipart + +import "net/textproto" + +type File interface { + Read(p []byte) (n int, err error) + Close() error +} + +type FileHeader struct { + Filename string + Header textproto.MIMEHeader + Size int64 +} + +func (fh *FileHeader) Open() (File, error) { return nil, nil } + +type Form struct { + Value map[string][]string + File map[string][]*FileHeader +} +`, + + "net/http": `package http + +import ( + "context" + "io" + "mime/multipart" + "net/url" +) + +type Header map[string][]string + +func (h Header) Get(key string) string { return "" } + +func (h Header) Set(key, value string) {} + +type Request struct { + Method string + URL *url.URL + Header Header + Body io.ReadCloser + Form url.Values + PostForm url.Values + MultipartForm *multipart.Form + Pattern string +} + +func (r *Request) Context() context.Context { return nil } + +func (r *Request) PathValue(name string) string { return "" } + +func (r *Request) FormValue(key string) string { return "" } + +func (r *Request) ParseForm() error { return nil } + +func (r *Request) ParseMultipartForm(maxMemory int64) error { return nil } + +type ResponseWriter interface { + Header() Header + Write([]byte) (int, error) + WriteHeader(statusCode int) +} + +type Flusher interface { + Flush() +} + +type Handler interface { + ServeHTTP(ResponseWriter, *Request) +} + +type HandlerFunc func(ResponseWriter, *Request) + +func (f HandlerFunc) ServeHTTP(w ResponseWriter, r *Request) {} + +type ServeMux struct{} + +func (mux *ServeMux) Handle(pattern string, handler Handler) {} + +func (mux *ServeMux) HandleFunc(pattern string, handler func(ResponseWriter, *Request)) {} +`, + + "net/textproto": `package textproto + +type MIMEHeader map[string][]string +`, + + "net/url": `package url + +type URL struct { + Scheme string + Host string + Path string + RawQuery string +} + +type Values map[string][]string + +func (v Values) Get(key string) string { return "" } + +func PathEscape(s string) string { return "" } +`, + + "time": `package time + +type Duration int64 + +type Time struct { + wall uint64 + ext int64 +} + +func (t Time) MarshalText() ([]byte, error) { return nil, nil } + +func (t *Time) UnmarshalText(data []byte) error { return nil } + +func Now() Time { return Time{} } +`, +} diff --git a/internal/typestest/typestest.go b/internal/typestest/typestest.go new file mode 100644 index 00000000..a8073c0d --- /dev/null +++ b/internal/typestest/typestest.go @@ -0,0 +1,198 @@ +// Package typestest type checks Go source in memory, for tests that need +// go/types values without a module on disk. +// +// Loading a real package runs the go command, and type checking the +// standard library from source takes seconds: net/http alone pulls in +// most of it. Route resolution and generation only ever ask a handful of +// questions of the standard library -- is this parameter an +// *http.Request, does this type implement encoding.TextUnmarshaler -- so +// the packages here are stubs declaring just the API muxt reads, with +// the real import paths and names. Checking them takes microseconds. +// +// A stub declares an identifier with the signature the real package +// gives it, or not at all. Add to one when a test needs more of it. +package typestest + +import ( + "fmt" + "go/ast" + "go/parser" + "go/token" + "go/types" + "slices" + "sort" + "sync" + "testing" +) + +// FileSet positions the stub packages and every package Check type +// checks. It is safe for concurrent use. +var FileSet = token.NewFileSet() + +// Check type checks files as the package with the import path path. Each +// entry maps a file name to its source. Imports resolve to the stub +// standard library packages; importing anything else is an error. +func Check(path string, files map[string]string) (*types.Package, error) { + checked, err := CheckSyntax(path, files) + if err != nil { + return nil, err + } + return checked.Types, nil +} + +// Checked is a package type checked by CheckSyntax, with the syntax and +// type information a test walks to find what the source does. +type Checked struct { + Types *types.Package + Syntax []*ast.File + Info *types.Info +} + +// CheckSyntax is Check, also returning the parsed files, in file name +// order, and the type information recorded for them. +func CheckSyntax(path string, files map[string]string) (*Checked, error) { + std, err := stdlib() + if err != nil { + return nil, err + } + info := &types.Info{ + Types: make(map[ast.Expr]types.TypeAndValue), + Instances: make(map[*ast.Ident]types.Instance), + Defs: make(map[*ast.Ident]types.Object), + Uses: make(map[*ast.Ident]types.Object), + Implicits: make(map[ast.Node]types.Object), + Selections: make(map[*ast.SelectorExpr]*types.Selection), + Scopes: make(map[ast.Node]*types.Scope), + } + pkg, syntax, err := checkWithInfo(path, files, importerFunc(func(p string) (*types.Package, error) { + if pkg, ok := std[p]; ok { + return pkg, nil + } + return nil, fmt.Errorf("typestest has no stub for package %q", p) + }), info) + if err != nil { + return nil, err + } + return &Checked{Types: pkg, Syntax: syntax, Info: info}, nil +} + +// MustCheck is Check for a single file, failing t when the source does +// not type check. +func MustCheck(t testing.TB, path, src string) *types.Package { + t.Helper() + pkg, err := Check(path, map[string]string{"source.go": src}) + if err != nil { + t.Fatal(err) + } + return pkg +} + +// Lookup finds a stub standard library package by import path. +func Lookup(path string) (*types.Package, bool) { + std, err := stdlib() + if err != nil { + panic(err) + } + pkg, ok := std[path] + return pkg, ok +} + +// Paths lists the import paths of the stub packages. +func Paths() []string { + paths := make([]string, 0, len(stubs)) + for path := range stubs { + paths = append(paths, path) + } + slices.Sort(paths) + return paths +} + +// Type returns the type the stub package at path declares as name, +// panicking when there is none: a test naming a type is asserting the +// stub declares it. +func Type(path, name string) types.Type { + pkg, ok := Lookup(path) + if !ok { + panic(fmt.Sprintf("typestest has no stub for package %q", path)) + } + obj := pkg.Scope().Lookup(name) + if obj == nil { + panic(fmt.Sprintf("typestest stub %q does not declare %s", path, name)) + } + return obj.Type() +} + +type importerFunc func(path string) (*types.Package, error) + +func (fn importerFunc) Import(path string) (*types.Package, error) { return fn(path) } + +func check(path string, files map[string]string, importer types.Importer) (*types.Package, error) { + pkg, _, err := checkWithInfo(path, files, importer, nil) + return pkg, err +} + +func checkWithInfo(path string, files map[string]string, importer types.Importer, info *types.Info) (*types.Package, []*ast.File, error) { + names := make([]string, 0, len(files)) + for name := range files { + names = append(names, name) + } + sort.Strings(names) + syntax := make([]*ast.File, 0, len(files)) + for _, name := range names { + file, err := parser.ParseFile(FileSet, name, files[name], parser.ParseComments|parser.SkipObjectResolution) + if err != nil { + return nil, nil, err + } + syntax = append(syntax, file) + } + config := types.Config{Importer: importer} + pkg, err := config.Check(path, FileSet, syntax, info) + return pkg, syntax, err +} + +var stdlib = sync.OnceValues(func() (map[string]*types.Package, error) { + checked := make(map[string]*types.Package, len(stubs)) + var importing []string + var importPath func(path string) (*types.Package, error) + importPath = func(path string) (*types.Package, error) { + if pkg, ok := checked[path]; ok { + return pkg, nil + } + src, ok := stubs[path] + if !ok { + return nil, fmt.Errorf("typestest stub imports %q, which has no stub", path) + } + if slices.Contains(importing, path) { + return nil, fmt.Errorf("typestest stubs import each other in a cycle: %v", append(importing, path)) + } + importing = append(importing, path) + defer func() { importing = importing[:len(importing)-1] }() + pkg, err := check(path, map[string]string{path + ".go": src}, importerFunc(importPath)) + if err != nil { + return nil, fmt.Errorf("typestest stub %q: %w", path, err) + } + checked[path] = pkg + return pkg, nil + } + for _, path := range Paths() { + if _, err := importPath(path); err != nil { + return nil, err + } + } + return checked, nil +}) + +// Packages returns the stub standard library packages by import path, the +// shape of source.Package's Imports. The map is a copy; the packages are +// shared. +func Packages() map[string]*types.Package { + std, err := stdlib() + if err != nil { + panic(err) + } + packages := make(map[string]*types.Package, len(std)) + for path, pkg := range std { + packages[path] = pkg + } + return packages +} diff --git a/internal/typestest/typestest_test.go b/internal/typestest/typestest_test.go new file mode 100644 index 00000000..6cbf264f --- /dev/null +++ b/internal/typestest/typestest_test.go @@ -0,0 +1,53 @@ +package typestest_test + +import ( + "go/types" + "strings" + "testing" + + "github.com/typelate/muxt/internal/typestest" +) + +func TestStubsTypeCheck(t *testing.T) { + for _, path := range typestest.Paths() { + if _, ok := typestest.Lookup(path); !ok { + t.Errorf("Lookup(%q) found no package", path) + } + } +} + +func TestCheck(t *testing.T) { + t.Run("imports resolve to the stubs", func(t *testing.T) { + pkg := typestest.MustCheck(t, "example.com/server", `package server + +import ( + "context" + "net/http" + "time" +) + +type Server struct{} + +func (Server) Get(ctx context.Context, request *http.Request, at time.Time) string { return "" } +`) + obj, _, _ := types.LookupFieldOrMethod(pkg.Scope().Lookup("Server").Type(), true, pkg, "Get") + sig := obj.Type().(*types.Signature) + if got := sig.Params().At(1).Type(); !types.Identical(got, types.NewPointer(typestest.Type("net/http", "Request"))) { + t.Errorf("request parameter has type %s, want the stub *http.Request", got) + } + }) + t.Run("the stub time.Time is a TextUnmarshaler", func(t *testing.T) { + unmarshaler := typestest.Type("encoding", "TextUnmarshaler").Underlying().(*types.Interface) + if !types.Implements(types.NewPointer(typestest.Type("time", "Time")), unmarshaler) { + t.Error("*time.Time does not implement encoding.TextUnmarshaler") + } + }) + t.Run("a package with no stub is an error", func(t *testing.T) { + _, err := typestest.Check("example.com/server", map[string]string{ + "source.go": "package server\n\nimport _ \"os\"\n", + }) + if err == nil || !strings.Contains(err.Error(), `no stub for package "os"`) { + t.Errorf("Check error = %v, want one naming the missing stub", err) + } + }) +}