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)
+ }
+ })
+}