diff --git a/CLAUDE.md b/CLAUDE.md index 2963ea2c..7d687533 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -25,15 +25,32 @@ source.Package: types + templates variables ./internal/source ↓ muxt.Definitions, muxt.ResolveCall Resolved routes (muxt.Definition) ./internal/muxt ↓ -Generated files / check reports ./internal/generate, ./internal/analysis +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. Route resolution, generation and the -template checks read only that, so they can be handed values built in memory. +definitions and ExecuteTemplate calls. Everything below it reads only that; +the mutation run, which loads again for `--diff` and with test files, loads +through `internal/load` too. + +**The standard library is asked, not copied.** Route resolution asks a +`muxt.Checker` what a reserved argument binds to and which types marshal to +and from text. `load.StandardLibrary` is the one implementation, answering from +the official standard library a run loaded. Tests of resolution and generation +use `internal/muxt/muxttest`, which builds the counterfeiter fake in +`internal/muxt/muxtfakes` over stand-in types declared in the test's own +source (`muxttest.StandInChecker(t, pkg)`, or `muxttest.NewChecker()` with +`Binds`, `ParsesFromText` and `FormatsAsText`), so they state muxt's rules +rather than one library version's shape. Regenerate the fake with +`go generate ./internal/muxt`. Tests that need real types use +`internal/load/loadtest`, which type checks package source against the +official standard library's export data without loading the package graph. +It loads one package whose imports are all in the standard library, embeds +only its top-level files, and has no test variants; behavior that depends on +more stays in `cmd/muxt` scripts. **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 @@ -93,41 +110,81 @@ Update the code in order: 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. + (`generate-fake-server` and `explore-module` do not go through it yet.) +- **What a command does with a valid configuration:** + `internal/{generate,analysis,mutation}/testdata//*.txtar` snapshot + generated files, check reports, the route and template listings, and + mutation dry runs, from packages loaded in memory in milliseconds. Each + archive is self-contained: the directory it is in names the command, its + `config.json` holds the configuration the command line in its header parses + into, and its `want/` files are what that produced. Run one with + `go test ./internal/analysis -run TestSnapshots/list-template-calls/calls`, + and rewrite them with + `go test ./internal/{generate,analysis,mutation} -run TestSnapshots -update`, + then review the diff. +- **Route names and call resolution:** `internal/muxt` tests type check source + in memory and resolve against a checker from `muxttest`. + +To add a snapshot case, write the archive in the command's directory: a +`config.json` holding what the command line parses into +(`internal/cli/configurations_test.go` states that parse), the package's files, +and no `want/` files. Then run the package's `TestSnapshots -update` and read +what it wrote. + +Integration scripts are for what needs the go command: generated code +compiling and serving requests, and files on disk. + ### 5. Verify Your Changes ```bash # Check for build/type errors go test ./cmd/muxt -# Run the formatter +# Run the formatters. goimports -local keeps this module's imports in +# their own group, after the third-party one; gofumpt does not group. go fmt ./... gofumpt -w . +goimports -local github.com/typelate/muxt -w . ``` ## Common Tasks ### Adding a New Feature -1. Create a test file: `cmd/muxt/testdata/reference_my_feature.txt` -2. Define the expected input (template) and output (generated code) -3. Run the test to see it fail -4. Update `internal/muxt/` generator functions -5. Run `go test ./cmd/muxt` until it passes +1. State it at the layer that owns it (see step 4 above): a snapshot archive + in `internal/generate/testdata/` for what is generated, + `internal/analysis/testdata/` for what is reported, + `internal/mutation/testdata/` for what a run mutates, a test in + `internal/muxt/` for how a name resolves, and a flag case in + `internal/cli/configurations_test.go` for a new flag +2. Run the test to see it fail +3. Update the package that owns the behavior: `internal/muxt/`, + `internal/generate/`, `internal/analysis/` or `internal/mutation/` +4. Rewrite the snapshot with `-update` and review the diff +5. Add `cmd/muxt/testdata/reference_my_feature.txt` when the generated code + must compile and serve requests, and run `go test ./cmd/muxt` ### Fixing a Bug -1. Create a test file: `cmd/muxt/testdata/err_bug_description.txt` or update an existing test -2. Reproduce the bug in the test -3. Run `go test ./cmd/muxt` to confirm failure -4. Fix the bug in `internal/muxt/` -5. Run `go test ./cmd/muxt` to confirm the fix +1. Reproduce it in the lowest test that can: a snapshot archive, an + `internal/muxt/` test, or, when it needs the go command, + `cmd/muxt/testdata/err_bug_description.txt` +2. Run the test to confirm failure +3. Fix the bug +4. Run the test, then `go test ./...`, to confirm the fix ### Adding Error Detection -1. Create a test: `cmd/muxt/testdata/err_error_name.txt` -2. Define input that should produce an error -3. Add validation logic to `internal/muxt/` -4. Verify the error message is clear +1. Add an `err_` snapshot archive (in `internal/generate/testdata/`, + `internal/analysis/testdata/` or `internal/mutation/testdata/`) whose + `want/error.txt` is the message +2. Add the validation to the package that reports it: `internal/muxt/` for a + route name, otherwise the package the archive belongs to +3. Verify the error message is clear ### Improving Documentation @@ -161,6 +218,9 @@ ls cmd/muxt/testdata/err_*.txt - `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/muxt/muxtfakes/` — The counterfeiter fake of `muxt.Checker`, generated by `go generate ./internal/muxt` +- `internal/muxt/muxttest/` — Builds that fake from what a test says the standard library looks like, plus import-free type checking +- `internal/load/loadtest/` — A loaded package type checked against the official standard library, for tests that go through `internal/load` - `internal/cli/` — Command-line interface - `cmd/muxt/` — Command entry point @@ -229,7 +289,8 @@ go -C ./cmd/muxt/testdata/debug-test test -v ## Pull Request Checklist - [ ] Tests pass: `go test ./...` -- [ ] Code formatted: `go fmt ./...` and `gofumpt -w .` +- [ ] Code formatted: `go fmt ./...`, `gofumpt -w .` and + `goimports -local github.com/typelate/muxt -w .` - [ ] New features have test files with clear naming - [ ] Error conditions are documented with `err_*` tests - [ ] No unnecessary changes to generated output diff --git a/go.mod b/go.mod index 5521b282..d82cff18 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/typelate/muxt -go 1.26.0 +go 1.27.0 require ( github.com/dustin/go-humanize v1.1.0 diff --git a/internal/analysis/configuration_test.go b/internal/analysis/configuration_test.go new file mode 100644 index 00000000..e7148ae8 --- /dev/null +++ b/internal/analysis/configuration_test.go @@ -0,0 +1,58 @@ +package analysis_test + +import ( + "encoding/json/v2" + "regexp" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/configjson" +) + +// TestListingConfigurationJSON states how a listing's configuration reads +// and writes as JSON, which is how an archive in testdata holds the one it +// runs with: a --match pattern is the text it was written as, and a field +// the command line left alone is null rather than an empty list, so a +// configuration read back is the one the command line produced. +// regexp.Regexp reads and writes itself, as a TextMarshaler; configjson +// says the rest. +func TestListingConfigurationJSON(t *testing.T) { + for _, tt := range []struct { + name string + config analysis.TemplateCallersConfiguration + want string + }{ + { + name: "no patterns", + config: analysis.TemplateCallersConfiguration{TemplatesVariables: []string{"templates"}}, + want: `{"TemplatesVariables":["templates"],"FilterTemplates":null}`, + }, + { + name: "a pattern is the text it was written as", + config: analysis.TemplateCallersConfiguration{TemplatesVariables: []string{"pages"}, FilterTemplates: []*regexp.Regexp{regexp.MustCompile("^head")}}, + want: `{"TemplatesVariables":["pages"],"FilterTemplates":["^head"]}`, + }, + } { + t.Run(tt.name, func(t *testing.T) { + written, err := json.Marshal(tt.config, configjson.Options()) + require.NoError(t, err) + assert.JSONEq(t, tt.want, string(written)) + + var read analysis.TemplateCallersConfiguration + require.NoError(t, json.Unmarshal(written, &read, configjson.Options())) + assert.Equal(t, tt.config, read) + }) + } +} + +// TestListingConfigurationJSONRejectsABadPattern states that a pattern +// that does not compile is reported where it was read. +func TestListingConfigurationJSONRejectsABadPattern(t *testing.T) { + var config analysis.TemplateCallsConfiguration + err := json.Unmarshal([]byte(`{"FilterTemplates":["("]}`), &config, configjson.Options()) + require.ErrorContains(t, err, "error parsing regexp") + require.ErrorContains(t, err, "FilterTemplates") +} diff --git a/internal/analysis/module.go b/internal/analysis/module.go index c7d11e1f..d4615a92 100644 --- a/internal/analysis/module.go +++ b/internal/analysis/module.go @@ -3,13 +3,14 @@ package analysis import ( "bufio" "bytes" - "encoding/json" + "encoding/json/v2" "io" + "maps" "os" "os/exec" "path/filepath" "regexp" - "sort" + "slices" "strings" "github.com/spf13/pflag" @@ -38,11 +39,11 @@ type PackageConfig struct { ReceiverType string `json:"receiverType,omitempty"` ReceiverPackage string `json:"receiverPackage,omitempty"` TemplateRoutePathsType string `json:"templateRoutePathsType"` - OutputHTMX bool `json:"outputHTMX,omitempty"` - OutputDatastar bool `json:"outputDatastar,omitempty"` - Logger bool `json:"logger,omitempty"` - PathPrefix bool `json:"pathPrefix,omitempty"` - Middleware bool `json:"middleware,omitempty"` + OutputHTMX bool `json:"outputHTMX,omitzero"` + OutputDatastar bool `json:"outputDatastar,omitzero"` + Logger bool `json:"logger,omitzero"` + PathPrefix bool `json:"pathPrefix,omitzero"` + Middleware bool `json:"middleware,omitzero"` } type PackageCommands struct { @@ -155,14 +156,8 @@ func NewModule(workingDirectory string, addFlags func(*pflag.FlagSet, *generate. return nil, err } - dirs := make([]string, 0, len(dirMap)) - for dir := range dirMap { - dirs = append(dirs, dir) - } - sort.Strings(dirs) - var packages []PackageInfo - for _, dir := range dirs { + for _, dir := range slices.Sorted(maps.Keys(dirMap)) { entry := dirMap[dir] var config generate.RoutesFileConfiguration set := pflag.NewFlagSet("parse-header", pflag.ContinueOnError) diff --git a/internal/analysis/snapshot_test.go b/internal/analysis/snapshot_test.go new file mode 100644 index 00000000..5559a701 --- /dev/null +++ b/internal/analysis/snapshot_test.go @@ -0,0 +1,262 @@ +package analysis_test + +import ( + "bytes" + "encoding/json/v2" + "errors" + "flag" + "fmt" + "io" + "log" + "os" + "path/filepath" + "reflect" + "slices" + "strings" + "testing" + + "golang.org/x/tools/txtar" + + "github.com/typelate/muxt/internal/analysis" + "github.com/typelate/muxt/internal/configjson" + "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's directory names the command it runs -- check, list-routes, +// list-template-callers or list-template-calls -- so one case runs with +// -run TestSnapshots/list-template-calls/calls. The archive holds the rest +// of the case: the configuration it runs with, its inputs, and what it +// reports. +// +// - config.json is the configuration, as the command line named in the +// archive's header parses into it. TestCommandLineConfigurations in +// internal/cli states that parse; this states what the command does +// with the result. +// - Go files and .gohtml files are loaded as example.com/server by +// internal/load/loadtest -- type checked against the official standard +// library, without loading the package graph -- 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 what a listing +// writes, want/checked.txt for the number of ExecuteTemplate calls +// check returns, 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) { + if stray, _ := filepath.Glob(filepath.Join("testdata", "*.txtar")); len(stray) > 0 { + t.Fatalf("%s is not in a command's directory", stray[0]) + } + directories, err := os.ReadDir("testdata") + if err != nil { + t.Fatal(err) + } + for _, directory := range directories { + if !directory.IsDir() { + continue + } + command := directory.Name() + newConfiguration, ok := commands[command] + if !ok { + t.Errorf("testdata/%s names no command this package snapshots", command) + continue + } + archives, err := filepath.Glob(filepath.Join("testdata", command, "*.txtar")) + if err != nil { + t.Fatal(err) + } + t.Run(command, func(t *testing.T) { + for _, archivePath := range archives { + runSnapshot(t, archivePath, newConfiguration) + } + }) + } +} + +// commands are the directories in testdata, each named for the command its +// archives run, with the configuration that command's config.json holds. +var commands = map[string]func() any{ + "check": func() any { return new(analysis.CheckConfiguration) }, + "list-routes": func() any { return new(analysis.DefinitionsConfiguration) }, + "list-template-callers": func() any { return new(analysis.TemplateCallersConfiguration) }, + "list-template-calls": func() any { return new(analysis.TemplateCallsConfiguration) }, +} + +// runSnapshot compares one archive's want/ files with what its command +// reports, or with -update rewrites them. +func runSnapshot(t *testing.T, archivePath string, newConfiguration func() any) { + t.Helper() + t.Run(strings.TrimSuffix(filepath.Base(archivePath), ".txtar"), func(t *testing.T) { + archive, err := txtar.ParseFile(archivePath) + if err != nil { + t.Fatal(err) + } + got := snapshot(t, configuration(t, archive, newConfiguration()), 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]) + } + } + }) +} + +// configuration reads the archive's config.json into config, a pointer to +// the configuration its directory's command holds, and returns the +// configuration by value, as a command holds it. +func configuration(t *testing.T, archive *txtar.Archive, config any) any { + t.Helper() + for _, file := range archive.Files { + if file.Name != "config.json" { + continue + } + if err := json.Unmarshal(file.Data, config, configjson.Options()); err != nil { + t.Fatalf("config.json: %v", err) + } + return reflect.ValueOf(config).Elem().Interface() + } + t.Fatal("the archive has no config.json") + return nil +} + +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/") || file.Name == "config.json" { + 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) + got["checked.txt"] = fmt.Sprintf("%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/testdata/check/check_bad_route_name.txtar b/internal/analysis/testdata/check/check_bad_route_name.txtar new file mode 100644 index 00000000..23a1c30b --- /dev/null +++ b/internal/analysis/testdata/check/check_bad_route_name.txtar @@ -0,0 +1,42 @@ +A malformed route template name is reported with its position. + +Command line: muxt check -v +-- config.json -- +{ + "Verbose": true, + "TemplatesVariables": [ + "templates" + ] +} +-- 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/checked.txt -- +1 +-- 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 diff --git a/internal/analysis/testdata/check/check_passes.txtar b/internal/analysis/testdata/check/check_passes.txtar new file mode 100644 index 00000000..d79a4d88 --- /dev/null +++ b/internal/analysis/testdata/check/check_passes.txtar @@ -0,0 +1,35 @@ +Each ExecuteTemplate call type checks the template it names, and the +templates it reaches through {{template}}. + +Command line: muxt check +-- config.json -- +{ + "Verbose": false, + "TemplatesVariables": [ + "templates" + ] +} +-- 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/checked.txt -- +1 diff --git a/internal/analysis/testdata/check/check_template_not_found.txtar b/internal/analysis/testdata/check/check_template_not_found.txtar new file mode 100644 index 00000000..fff00550 --- /dev/null +++ b/internal/analysis/testdata/check/check_template_not_found.txtar @@ -0,0 +1,39 @@ +An ExecuteTemplate call naming a template the set does not define. + +Command line: muxt check +-- config.json -- +{ + "Verbose": false, + "TemplatesVariables": [ + "templates" + ] +} +-- 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/checked.txt -- +1 +-- want/error.txt -- +1 error +-- want/log.txt -- +server.go:16:9 ExecuteTemplate "index" Page + - template "index" not found + diff --git a/internal/analysis/testdata/check/check_unused_templates.txtar b/internal/analysis/testdata/check/check_unused_templates.txtar new file mode 100644 index 00000000..db448aff --- /dev/null +++ b/internal/analysis/testdata/check/check_unused_templates.txtar @@ -0,0 +1,44 @@ +A template nothing executes is reported; a route template is reported +as waiting for muxt generate, and a partial only it uses is not reported. + +Command line: muxt check +-- config.json -- +{ + "Verbose": false, + "TemplatesVariables": [ + "templates" + ] +} +-- 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/checked.txt -- +1 +-- 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" diff --git a/internal/analysis/testdata/check/check_wrong_field.txtar b/internal/analysis/testdata/check/check_wrong_field.txtar new file mode 100644 index 00000000..5a3e3205 --- /dev/null +++ b/internal/analysis/testdata/check/check_wrong_field.txtar @@ -0,0 +1,53 @@ +A field the data type does not have is reported at the call and in the +template. + +Command line: muxt check +-- config.json -- +{ + "Verbose": false, + "TemplatesVariables": [ + "templates" + ] +} +-- 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/checked.txt -- +1 +-- 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 + } + diff --git a/internal/analysis/testdata/list-routes/routes.txtar b/internal/analysis/testdata/list-routes/routes.txtar new file mode 100644 index 00000000..b9886ce5 --- /dev/null +++ b/internal/analysis/testdata/list-routes/routes.txtar @@ -0,0 +1,41 @@ +The route listing shows each route template's name, and the receiver's +methods. + +Command line: muxt --use-receiver-type=T +-- config.json -- +{ + "Verbose": false, + "ReceiverPackage": "", + "ReceiverType": "T", + "TemplatesVariables": [ + "templates" + ] +} +-- 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/analysis/testdata/list-template-callers/callers.txtar b/internal/analysis/testdata/list-template-callers/callers.txtar new file mode 100644 index 00000000..4119f6b3 --- /dev/null +++ b/internal/analysis/testdata/list-template-callers/callers.txtar @@ -0,0 +1,57 @@ +Where each template is executed or called from, with the data type. + +Command line: muxt list-template-callers +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "FilterTemplates": null +} +-- 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/list-template-callers/callers_match.txtar b/internal/analysis/testdata/list-template-callers/callers_match.txtar new file mode 100644 index 00000000..e5a33ae8 --- /dev/null +++ b/internal/analysis/testdata/list-template-callers/callers_match.txtar @@ -0,0 +1,44 @@ +Where the matching templates are executed or called from. + +Command line: muxt list-template-callers --match=^head +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "FilterTemplates": [ + "^head" + ] +} +-- 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/list-template-calls/calls.txtar b/internal/analysis/testdata/list-template-calls/calls.txtar new file mode 100644 index 00000000..8ae72ea0 --- /dev/null +++ b/internal/analysis/testdata/list-template-calls/calls.txtar @@ -0,0 +1,47 @@ +The templates each template calls, with the data type. + +Command line: muxt list-template-calls +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "FilterTemplates": null +} +-- 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/astgen/gen.go b/internal/astgen/gen.go index 050ddd1c..d148ea1d 100644 --- a/internal/astgen/gen.go +++ b/internal/astgen/gen.go @@ -2,7 +2,6 @@ package astgen import ( "go/ast" - "go/types" ) // ImportManager interface abstracts the import management functionality @@ -16,9 +15,6 @@ type ImportManager interface { // ImportSpecs returns all registered import specs ImportSpecs() []*ast.ImportSpec - - // TypeASTExpression converts a types.Type to an AST expression - TypeASTExpression(tp types.Type) (ast.Expr, error) } // ExportedIdentifier creates a selector expression for an exported identifier diff --git a/internal/astgen/pkg.go b/internal/astgen/pkg.go index 4cd03f6f..1bcf2a0c 100644 --- a/internal/astgen/pkg.go +++ b/internal/astgen/pkg.go @@ -1,7 +1,8 @@ package astgen import ( - "encoding/json" + "encoding/json/jsontext" + "encoding/json/v2" "go/ast" "go/token" "go/types" @@ -24,8 +25,12 @@ func NewTypeFormatter(outputPkgPath string) *TypeFormatter { } } -func (tf *TypeFormatter) MarshalJSON() ([]byte, error) { - return json.MarshalIndent(tf.Imports, "", " ") +// MarshalJSONTo writes the imports the formatter collected. It writes +// through the encoder it is handed, so whatever the whole document is +// written with -- the indentation, and the key order a reproducible listing +// needs -- holds for the imports too. +func (tf *TypeFormatter) MarshalJSONTo(encoder *jsontext.Encoder) error { + return json.MarshalEncode(encoder, tf.Imports) } func (tf *TypeFormatter) Qualifier(pkg *types.Package) string { diff --git a/internal/astgen/strconv.go b/internal/astgen/strconv.go index f8f15506..8073fb7c 100644 --- a/internal/astgen/strconv.go +++ b/internal/astgen/strconv.go @@ -4,10 +4,16 @@ import ( "fmt" "go/ast" "go/types" + + "github.com/typelate/muxt/internal/source" ) -// ConvertToString converts a variable to its string representation based on its basic kind -func ConvertToString(im ImportManager, variable ast.Expr, kind types.BasicKind) (ast.Expr, error) { +// ConvertToString formats variable, whose type is tp, as a string. +func ConvertToString(im ImportManager, variable ast.Expr, tp source.Type) (ast.Expr, error) { + kind, ok := tp.Basic() + if !ok { + return nil, fmt.Errorf("unsupported type for path parameters") + } switch kind { case types.Bool, types.UntypedBool: return FormatBool(im, variable), nil diff --git a/internal/cli/commands.go b/internal/cli/commands.go index a5762afd..f65c5127 100644 --- a/internal/cli/commands.go +++ b/internal/cli/commands.go @@ -1,10 +1,10 @@ package cli import ( - "bytes" "cmp" _ "embed" - "encoding/json" + "encoding/json/jsontext" + "encoding/json/v2" "errors" "fmt" "go/token" @@ -44,6 +44,36 @@ 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 for these six commands it is decided +// before a package is loaded; --format alone is still read when a result is +// written. 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. generate-fake-server and explore-module load packages +// in their own RunE and have no runner yet. +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,26 +106,7 @@ func Commands(wd string, args []string, getEnv func(string) string, stdout, stde return err } cmd.SilenceUsage = true - _, pl, err := load.Packages(*workingDirectory, rootCommandConfig.ReceiverPackage) - if err != nil { - return err - } - pkg, receiver, err := load.RoutesSource(*workingDirectory, pl, rootCommandConfig) - 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 + return run.routes(cmd, *workingDirectory, rootCommandConfig) }, } rootCmd.PersistentFlags().StringVarP(&changeDir, "change-directory", "C", "", "change the working directory") @@ -111,14 +122,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 @@ -132,7 +143,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, @@ -153,26 +164,7 @@ func checkCommand(workingDirectory *string) *cobra.Command { } } cmd.SilenceUsage = true - _, pl, err := load.Packages(*workingDirectory) - if err != nil { - return err - } - logger := log.New(cmd.ErrOrStderr(), "", 0) - warnPartialAST(logger, pl) - pkg, err := load.Package(*workingDirectory, pl, config.TemplatesVariables) - if err != nil { - return checkFailure(cmd, err) - } - checked, err := analysis.Check(config, logger, pkg) - if err != nil { - return checkFailure(cmd, 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) }, } @@ -186,7 +178,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 @@ -247,13 +239,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) }, } @@ -299,7 +285,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 @@ -314,7 +300,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) @@ -345,106 +330,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 - _, pl, err := load.Packages(*workingDirectory, config.ReceiverPackage) - if err != nil { - return err - } - warnPartialAST(log.New(cmd.ErrOrStderr(), "", 0), pl) - pkg, receiver, err := load.GenerateSource(*workingDirectory, pl, config) - if err != nil { - printMultiLineError(cmd, err) - return err - } - files, err := generate.TemplateRoutesFiles(*workingDirectory, 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(*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) }, } @@ -524,18 +415,21 @@ 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 deprecatedTemplatesVar string @@ -559,19 +453,7 @@ func listTemplateCallersCommand(wd *string) *cobra.Command { config.FilterTemplates = append(config.FilterTemplates, pat) } - _, 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) + return run(cmd, *wd, config) }, } @@ -582,7 +464,7 @@ func listTemplateCallersCommand(wd *string) *cobra.Command { 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 patterns []string @@ -606,19 +488,7 @@ func listTemplateCallsCommand(wd *string) *cobra.Command { config.FilterTemplates = append(config.FilterTemplates, pat) } - _, 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) + return run(cmd, *wd, config) }, } @@ -1058,6 +928,19 @@ This command is intended for exploratory use only.`, return cmd } +// resultJSON writes a --format=json result as encoding/json wrote it before +// muxt moved to encoding/json/v2, so a script reading it sees no change: +// map members sorted by key, a nil list or map as null, and <, >, &, U+2028 +// and U+2029 escaped. +var resultJSON = json.JoinOptions( + jsontext.WithIndent("\t"), + json.Deterministic(true), + json.FormatNilSliceAsNull(true), + json.FormatNilMapAsNull(true), + jsontext.EscapeForHTML(true), + jsontext.EscapeForJS(true), +) + func writeResult(cmd *cobra.Command, w io.Writer, result io.WriterTo) error { format, err := cmd.Flags().GetString("format") if err != nil { @@ -1065,7 +948,7 @@ func writeResult(cmd *cobra.Command, w io.Writer, result io.WriterTo) error { } switch format { case "json": - buf, err := json.MarshalIndent(result, "", "\t") + buf, err := json.Marshal(result, resultJSON) 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..87fc3b09 --- /dev/null +++ b/internal/cli/configurations_test.go @@ -0,0 +1,556 @@ +package cli + +import ( + "io" + "path/filepath" + "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 test of what a +// command does with a configuration repeats one of them beside the command +// line it stands for, so searching for the literal finds both. Command +// lines that cannot work are in TestCommandLineRejections, so 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 every explicit name", + args: "generate --output-exported-default-identifiers=false --output-receiver-interface=Handlers --output-template-data-type=Data --output-sse-template-data-type=SSEData --output-template-route-paths-type=Paths", + want: generate.RoutesFileConfiguration{ + MuxtVersion: "v1.2.3", + PackageName: "main", + RoutesFunction: "templateRoutes", + ReceiverInterface: "Handlers", + TemplateDataType: "Data", + SSETemplateDataType: "SSEData", + TemplateRoutePathsTypeName: "Paths", + TemplatesVariables: []string{"templates"}, + OutputFileName: "template_routes.go", + OutputMuxtVersion: true, + }, + }, + { + name: "unexported default identifiers with a receiver type", + args: "generate --use-receiver-type=Server --output-exported-default-identifiers=false", + 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", + 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: "template calls matching a pattern", + args: "list-template-calls --match=^Index$", + want: analysis.TemplateCallsConfiguration{FilterTemplates: []*regexp.Regexp{regexp.MustCompile("^Index$")}, 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: "a mutation dry run, quietly", + args: "test-template-mutations --dry-run --seed=1", + want: mutation.Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Seed: 1, SeedSet: true, MaxCases: mutation.DefaultMaxCases, Workers: 1}, + }, + { + name: "a mutation dry run of the templates a pattern matches", + args: "test-template-mutations --dry-run --seed=1 -v --template-pattern=^footer$", + want: mutation.Configuration{TemplatesVariables: []string{"templates"}, TemplatePattern: regexp.MustCompile("^footer$"), Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: mutation.DefaultMaxCases, Workers: 1}, + }, + { + name: "a mutation dry run with an operand budget", + args: "test-template-mutations --dry-run --seed=1 -v --max-cases=2", + want: mutation.Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: 2, Workers: 1}, + }, + { + name: "a mutation dry run of what changed since a revision", + args: "test-template-mutations --dry-run --seed=1 -v --diff=main", + want: mutation.Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Verbose: true, Seed: 1, SeedSet: true, MaxCases: mutation.DefaultMaxCases, Workers: 1, Diff: "main"}, + }, + { + name: "a mutation dry run of another templates variable", + args: "test-template-mutations --dry-run --seed=1 --use-templates-variable=pages", + want: mutation.Configuration{TemplatesVariables: []string{"pages"}, Packages: []string{}, DryRun: 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) + }) + } +} + +// TestChangeDirectory states the working directory a command runs in: the +// one it was started in, joined with -C when -C is relative. +func TestChangeDirectory(t *testing.T) { + for _, tt := range []struct { + args, want string + }{ + {args: "generate", want: "/work"}, + {args: "-C sub generate", want: "/work/sub"}, + {args: "-C ../other check", want: "/other"}, + {args: "-C /abs check", want: "/abs"}, + } { + t.Run(tt.args, func(t *testing.T) { + var got string + record := func(wd string) error { + got = wd + return nil + } + err := commands("/work", strings.Fields(tt.args), func(string) string { return "" }, func() (string, bool) { return "v1.2.3", true }, io.Discard, io.Discard, runners{ + check: func(_ *cobra.Command, wd string, _ analysis.CheckConfiguration) error { return record(wd) }, + generate: func(_ *cobra.Command, wd string, _ generate.RoutesFileConfiguration) error { return record(wd) }, + }) + require.NoError(t, err) + assert.Equal(t, filepath.FromSlash(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 ): `(`"}, + {name: "a go test flag muxt needs for itself", args: "test-template-mutations -- -overlay=other.json", wantErr: "go test flag -overlay cannot be passed through: muxt uses -overlay to deliver each mutant"}, + } { + 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/result_json_test.go b/internal/cli/result_json_test.go new file mode 100644 index 00000000..fe3c2edd --- /dev/null +++ b/internal/cli/result_json_test.go @@ -0,0 +1,60 @@ +package cli + +import ( + "bytes" + "fmt" + "io" + "strings" + "testing" + + "github.com/spf13/cobra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/astgen" +) + +// jsonResult has the shapes a --format=json result is made of: a list and +// a map a command may leave nil, the import map a listing carries, and +// template source. +type jsonResult struct { + Names []string + Counts map[string]int + Imports *astgen.TypeFormatter + Source string +} + +func (jsonResult) WriteTo(io.Writer) (int64, error) { return 0, nil } + +// TestWriteResultJSON states that --format=json writes a result as +// encoding/json did before muxt moved to encoding/json/v2: map members, +// including a listing's imports, sorted by key on every run; a nil list or +// map as null; and <, >, & and U+2028 escaped. +func TestWriteResultJSON(t *testing.T) { + imports := astgen.NewTypeFormatter("example.com/server") + for i := range 16 { + imports.Imports[fmt.Sprintf("example.com/pkg%02d", 15-i)] = fmt.Sprintf("pkg%02d", 15-i) + } + result := jsonResult{Imports: imports, Source: "

    {{.A}} & {{.B}}

    
"} + + var want strings.Builder + want.WriteString("{\n\t\"Names\": null,\n\t\"Counts\": null,\n\t\"Imports\": {\n") + for i := range 16 { + separator := "," + if i == 15 { + separator = "" + } + fmt.Fprintf(&want, "\t\t\"example.com/pkg%02d\": \"pkg%02d\"%s\n", i, i, separator) + } + want.WriteString("\t},\n\t\"Source\": \"\\u003cp\\u003e{{.A}} \\u0026 {{.B}}\\u003c/p\\u003e\\u2028\"\n}\n") + + cmd := &cobra.Command{} + cmd.Flags().String("format", "json", "") + // Map iteration order changes run to run, so one lucky ordering must not + // pass. + for range 20 { + var got bytes.Buffer + require.NoError(t, writeResult(cmd, &got, result)) + assert.Equal(t, want.String(), got.String()) + } +} diff --git a/internal/cli/run.go b/internal/cli/run.go new file mode 100644 index 00000000..84e4947f --- /dev/null +++ b/internal/cli/run.go @@ -0,0 +1,216 @@ +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" + "github.com/typelate/muxt/internal/muxt" +) + +// This file holds what each command does with its configuration: load the +// package, run the implementation, and write what it produced. What the +// flags decide, but for --format, 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) + pkg, err := load.Package(wd, pl, config.TemplatesVariables) + if err != nil { + return checkFailure(cmd, err) + } + checked, err := analysis.Check(config, logger, pkg) + if err != nil { + return checkFailure(cmd, 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 + } + defs, err := muxt.ResolveDefinitions(pkg, receiver, load.StandardLibrary(pl)) + if err != nil { + printMultiLineError(cmd, err) + return err + } + files, err := generate.TemplateRoutesFiles(wd, config, pkg, defs, 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/configjson/configjson.go b/internal/configjson/configjson.go new file mode 100644 index 00000000..2ed664f7 --- /dev/null +++ b/internal/configjson/configjson.go @@ -0,0 +1,23 @@ +// Package configjson says how muxt reads and writes a command's +// configuration as JSON. +// +// It is how the snapshot archives in internal/{generate,analysis,mutation} +// hold the configuration they run with, so a configuration read back is the +// one a command line produced: a field a command line left alone is null +// rather than an empty list, and a member the configuration does not +// declare is an error rather than a typo nothing reports. +// +// Patterns say nothing here: regexp.Regexp reads and writes itself as the +// text it was compiled from. +package configjson + +import "encoding/json/v2" + +// Options reads and writes a configuration the way muxt holds one. +func Options() json.Options { + return json.JoinOptions( + json.FormatNilSliceAsNull(true), + json.FormatNilMapAsNull(true), + json.RejectUnknownMembers(true), + ) +} diff --git a/internal/fake/check.go b/internal/fake/check.go new file mode 100644 index 00000000..b03c63a8 --- /dev/null +++ b/internal/fake/check.go @@ -0,0 +1,190 @@ +package fake + +import ( + "fmt" + "go/ast" + "go/parser" + "go/token" + "go/types" + "maps" + "slices" + "testing" + + "github.com/typelate/muxt/internal/muxt" +) + +// FileSet positions every package Check type checks. It is safe for +// concurrent use. +var FileSet = token.NewFileSet() + +// Check type checks files, by file name, as the package with import path +// path. The source may import nothing: a type it needs from outside the +// package is declared in it. +func Check(t testing.TB, path string, files map[string]string) *types.Package { + t.Helper() + names := slices.Sorted(maps.Keys(files)) + syntax := make([]*ast.File, 0, len(names)) + for _, name := range names { + file, err := parser.ParseFile(FileSet, name, files[name], parser.SkipObjectResolution) + if err != nil { + t.Fatal(err) + } + syntax = append(syntax, file) + } + config := types.Config{Importer: importerFunc(func(path string) (*types.Package, error) { + return nil, fmt.Errorf("muxttest source imports %q; declare a stand-in in the package instead", path) + })} + pkg, err := config.Check(path, FileSet, syntax, nil) + if err != nil { + t.Fatal(err) + } + return pkg +} + +type importerFunc func(path string) (*types.Package, error) + +func (fn importerFunc) Import(path string) (*types.Package, error) { return fn(path) } + +// CheckerBuilder says what a test wants the standard library to look +// like. Fake turns that into the generated muxt.Checker fake. +// +// Read a built checker as a sentence: the identifiers it binds, the types +// that parse from text, the type an upload binds to. +type CheckerBuilder struct { + scope map[string]types.Type + fileHeader types.Type + rawJSON types.Type + textUnmarshalers []types.Type + textMarshalers []types.Type +} + +// NewChecker starts a checker that answers nothing: every question a route +// asks it fails, which is what a route needing no request value expects. +func NewChecker() *CheckerBuilder { + return &CheckerBuilder{scope: make(map[string]types.Type)} +} + +// standIns are the stand-in type names StandInChecker binds each reserved +// argument identifier to, when the package declares them. +var standIns = []struct { + identifier string + typeName string + pointer bool +}{ + {identifier: muxt.TemplateNameScopeIdentifierHTTPRequest, typeName: "Request", pointer: true}, + {identifier: muxt.TemplateNameScopeIdentifierHTTPResponse, typeName: "ResponseWriter"}, + {identifier: muxt.TemplateNameScopeIdentifierContext, typeName: "Context"}, + {identifier: muxt.TemplateNameScopeIdentifierForm, typeName: "Values"}, + {identifier: muxt.TemplateNameScopeIdentifierMultipart, typeName: "Form", pointer: true}, + {identifier: muxt.TemplateNameScopeIdentifierRequestBody, typeName: "Reader"}, +} + +// StandInChecker binds each reserved argument identifier to the stand-in +// pkg declares for it -- request to *Request, ctx to Context, and so on -- +// and binds an upload to *FileHeader and a JSON body to RawMessage when +// they are declared. A type pkg does not declare is left unbound, so a +// route asking for it fails the way one asking for anything unbound does. +func StandInChecker(t testing.TB, pkg *types.Package) *CheckerBuilder { + t.Helper() + c := NewChecker() + for _, standIn := range standIns { + tp, ok := declared(pkg, standIn.typeName) + if !ok { + continue + } + if standIn.pointer { + tp = types.NewPointer(tp) + } + c.Binds(standIn.identifier, tp) + } + if tp, ok := declared(pkg, "FileHeader"); ok { + c.UploadsFilesAs(types.NewPointer(tp)) + } + if tp, ok := declared(pkg, "RawMessage"); ok { + c.DecodesJSONAs(tp) + } + return c +} + +// Binds says the reserved argument identifier binds to tp. +func (c *CheckerBuilder) Binds(identifier string, tp types.Type) *CheckerBuilder { + c.scope[identifier] = tp + return c +} + +// ParsesFromText says a pointer to each type implements +// encoding.TextUnmarshaler, so a request string parses into it. +func (c *CheckerBuilder) ParsesFromText(tps ...types.Type) *CheckerBuilder { + c.textUnmarshalers = append(c.textUnmarshalers, tps...) + return c +} + +// FormatsAsText says each type implements encoding.TextMarshaler, so a +// route path formats it as a path segment. +func (c *CheckerBuilder) FormatsAsText(tps ...types.Type) *CheckerBuilder { + c.textMarshalers = append(c.textMarshalers, tps...) + return c +} + +// UploadsFilesAs says tp is the multipart file header type. +func (c *CheckerBuilder) UploadsFilesAs(tp types.Type) *CheckerBuilder { + c.fileHeader = tp + return c +} + +// DecodesJSONAs says tp is the raw JSON message type. +func (c *CheckerBuilder) DecodesJSONAs(tp types.Type) *CheckerBuilder { + c.rawJSON = tp + return c +} + +// Fake returns the generated fake answering what the builder was told, and +// failing every other question the way a standard library that does not +// declare the type would. +func (c *CheckerBuilder) Fake() *Checker { + fake := new(Checker) + fake.ScopeTypeStub = func(identifier string) (types.Type, error) { + tp, ok := c.scope[identifier] + if !ok { + return nil, fmt.Errorf("the checker binds no type to %s", identifier) + } + return tp, nil + } + fake.FileHeaderStub = answer(c.fileHeader, "multipart file header") + fake.RawJSONStub = answer(c.rawJSON, "raw JSON message") + fake.TextUnmarshalerStub = func(tp types.Type) bool { return contains(c.textUnmarshalers, tp) } + fake.TextMarshalerStub = func(tp types.Type) bool { return contains(c.textMarshalers, tp) } + return fake +} + +func answer(tp types.Type, what string) func() (types.Type, error) { + return func() (types.Type, error) { + if tp == nil { + return nil, fmt.Errorf("the checker has no %s type", what) + } + return tp, nil + } +} + +func contains(list []types.Type, tp types.Type) bool { + return slices.ContainsFunc(list, func(candidate types.Type) bool { return types.Identical(candidate, tp) }) +} + +func declared(pkg *types.Package, name string) (types.Type, bool) { + obj := pkg.Scope().Lookup(name) + if obj == nil { + return nil, false + } + return obj.Type(), true +} + +// Lookup returns the type pkg declares as name, failing t when it declares +// none. +func Lookup(t testing.TB, pkg *types.Package, name string) types.Type { + t.Helper() + obj := pkg.Scope().Lookup(name) + if obj == nil { + t.Fatalf("%s declares no %s", pkg.Path(), name) + } + return obj.Type() +} diff --git a/internal/fake/checker.go b/internal/fake/checker.go new file mode 100644 index 00000000..a3f81dbd --- /dev/null +++ b/internal/fake/checker.go @@ -0,0 +1,395 @@ +// Code generated by counterfeiter. DO NOT EDIT. +package fake + +import ( + "go/types" + "sync" + + "github.com/typelate/muxt/internal/muxt" +) + +type Checker struct { + FileHeaderStub func() (types.Type, error) + fileHeaderMutex sync.RWMutex + fileHeaderArgsForCall []struct { + } + fileHeaderReturns struct { + result1 types.Type + result2 error + } + fileHeaderReturnsOnCall map[int]struct { + result1 types.Type + result2 error + } + RawJSONStub func() (types.Type, error) + rawJSONMutex sync.RWMutex + rawJSONArgsForCall []struct { + } + rawJSONReturns struct { + result1 types.Type + result2 error + } + rawJSONReturnsOnCall map[int]struct { + result1 types.Type + result2 error + } + ScopeTypeStub func(string) (types.Type, error) + scopeTypeMutex sync.RWMutex + scopeTypeArgsForCall []struct { + arg1 string + } + scopeTypeReturns struct { + result1 types.Type + result2 error + } + scopeTypeReturnsOnCall map[int]struct { + result1 types.Type + result2 error + } + TextMarshalerStub func(types.Type) bool + textMarshalerMutex sync.RWMutex + textMarshalerArgsForCall []struct { + arg1 types.Type + } + textMarshalerReturns struct { + result1 bool + } + textMarshalerReturnsOnCall map[int]struct { + result1 bool + } + TextUnmarshalerStub func(types.Type) bool + textUnmarshalerMutex sync.RWMutex + textUnmarshalerArgsForCall []struct { + arg1 types.Type + } + textUnmarshalerReturns struct { + result1 bool + } + textUnmarshalerReturnsOnCall map[int]struct { + result1 bool + } + invocations map[string][][]interface{} + invocationsMutex sync.RWMutex +} + +func (fake *Checker) FileHeader() (types.Type, error) { + fake.fileHeaderMutex.Lock() + ret, specificReturn := fake.fileHeaderReturnsOnCall[len(fake.fileHeaderArgsForCall)] + fake.fileHeaderArgsForCall = append(fake.fileHeaderArgsForCall, struct { + }{}) + stub := fake.FileHeaderStub + fakeReturns := fake.fileHeaderReturns + fake.recordInvocation("FileHeader", []interface{}{}) + fake.fileHeaderMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1, ret.result2 + } + return fakeReturns.result1, fakeReturns.result2 +} + +func (fake *Checker) FileHeaderCallCount() int { + fake.fileHeaderMutex.RLock() + defer fake.fileHeaderMutex.RUnlock() + return len(fake.fileHeaderArgsForCall) +} + +func (fake *Checker) FileHeaderCalls(stub func() (types.Type, error)) { + fake.fileHeaderMutex.Lock() + defer fake.fileHeaderMutex.Unlock() + fake.FileHeaderStub = stub +} + +func (fake *Checker) FileHeaderReturns(result1 types.Type, result2 error) { + fake.fileHeaderMutex.Lock() + defer fake.fileHeaderMutex.Unlock() + fake.FileHeaderStub = nil + fake.fileHeaderReturns = struct { + result1 types.Type + result2 error + }{result1, result2} +} + +func (fake *Checker) FileHeaderReturnsOnCall(i int, result1 types.Type, result2 error) { + fake.fileHeaderMutex.Lock() + defer fake.fileHeaderMutex.Unlock() + fake.FileHeaderStub = nil + if fake.fileHeaderReturnsOnCall == nil { + fake.fileHeaderReturnsOnCall = make(map[int]struct { + result1 types.Type + result2 error + }) + } + fake.fileHeaderReturnsOnCall[i] = struct { + result1 types.Type + result2 error + }{result1, result2} +} + +func (fake *Checker) RawJSON() (types.Type, error) { + fake.rawJSONMutex.Lock() + ret, specificReturn := fake.rawJSONReturnsOnCall[len(fake.rawJSONArgsForCall)] + fake.rawJSONArgsForCall = append(fake.rawJSONArgsForCall, struct { + }{}) + stub := fake.RawJSONStub + fakeReturns := fake.rawJSONReturns + fake.recordInvocation("RawJSON", []interface{}{}) + fake.rawJSONMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1, ret.result2 + } + return fakeReturns.result1, fakeReturns.result2 +} + +func (fake *Checker) RawJSONCallCount() int { + fake.rawJSONMutex.RLock() + defer fake.rawJSONMutex.RUnlock() + return len(fake.rawJSONArgsForCall) +} + +func (fake *Checker) RawJSONCalls(stub func() (types.Type, error)) { + fake.rawJSONMutex.Lock() + defer fake.rawJSONMutex.Unlock() + fake.RawJSONStub = stub +} + +func (fake *Checker) RawJSONReturns(result1 types.Type, result2 error) { + fake.rawJSONMutex.Lock() + defer fake.rawJSONMutex.Unlock() + fake.RawJSONStub = nil + fake.rawJSONReturns = struct { + result1 types.Type + result2 error + }{result1, result2} +} + +func (fake *Checker) RawJSONReturnsOnCall(i int, result1 types.Type, result2 error) { + fake.rawJSONMutex.Lock() + defer fake.rawJSONMutex.Unlock() + fake.RawJSONStub = nil + if fake.rawJSONReturnsOnCall == nil { + fake.rawJSONReturnsOnCall = make(map[int]struct { + result1 types.Type + result2 error + }) + } + fake.rawJSONReturnsOnCall[i] = struct { + result1 types.Type + result2 error + }{result1, result2} +} + +func (fake *Checker) ScopeType(arg1 string) (types.Type, error) { + fake.scopeTypeMutex.Lock() + ret, specificReturn := fake.scopeTypeReturnsOnCall[len(fake.scopeTypeArgsForCall)] + fake.scopeTypeArgsForCall = append(fake.scopeTypeArgsForCall, struct { + arg1 string + }{arg1}) + stub := fake.ScopeTypeStub + fakeReturns := fake.scopeTypeReturns + fake.recordInvocation("ScopeType", []interface{}{arg1}) + fake.scopeTypeMutex.Unlock() + if stub != nil { + return stub(arg1) + } + if specificReturn { + return ret.result1, ret.result2 + } + return fakeReturns.result1, fakeReturns.result2 +} + +func (fake *Checker) ScopeTypeCallCount() int { + fake.scopeTypeMutex.RLock() + defer fake.scopeTypeMutex.RUnlock() + return len(fake.scopeTypeArgsForCall) +} + +func (fake *Checker) ScopeTypeCalls(stub func(string) (types.Type, error)) { + fake.scopeTypeMutex.Lock() + defer fake.scopeTypeMutex.Unlock() + fake.ScopeTypeStub = stub +} + +func (fake *Checker) ScopeTypeArgsForCall(i int) string { + fake.scopeTypeMutex.RLock() + defer fake.scopeTypeMutex.RUnlock() + argsForCall := fake.scopeTypeArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *Checker) ScopeTypeReturns(result1 types.Type, result2 error) { + fake.scopeTypeMutex.Lock() + defer fake.scopeTypeMutex.Unlock() + fake.ScopeTypeStub = nil + fake.scopeTypeReturns = struct { + result1 types.Type + result2 error + }{result1, result2} +} + +func (fake *Checker) ScopeTypeReturnsOnCall(i int, result1 types.Type, result2 error) { + fake.scopeTypeMutex.Lock() + defer fake.scopeTypeMutex.Unlock() + fake.ScopeTypeStub = nil + if fake.scopeTypeReturnsOnCall == nil { + fake.scopeTypeReturnsOnCall = make(map[int]struct { + result1 types.Type + result2 error + }) + } + fake.scopeTypeReturnsOnCall[i] = struct { + result1 types.Type + result2 error + }{result1, result2} +} + +func (fake *Checker) TextMarshaler(arg1 types.Type) bool { + fake.textMarshalerMutex.Lock() + ret, specificReturn := fake.textMarshalerReturnsOnCall[len(fake.textMarshalerArgsForCall)] + fake.textMarshalerArgsForCall = append(fake.textMarshalerArgsForCall, struct { + arg1 types.Type + }{arg1}) + stub := fake.TextMarshalerStub + fakeReturns := fake.textMarshalerReturns + fake.recordInvocation("TextMarshaler", []interface{}{arg1}) + fake.textMarshalerMutex.Unlock() + if stub != nil { + return stub(arg1) + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *Checker) TextMarshalerCallCount() int { + fake.textMarshalerMutex.RLock() + defer fake.textMarshalerMutex.RUnlock() + return len(fake.textMarshalerArgsForCall) +} + +func (fake *Checker) TextMarshalerCalls(stub func(types.Type) bool) { + fake.textMarshalerMutex.Lock() + defer fake.textMarshalerMutex.Unlock() + fake.TextMarshalerStub = stub +} + +func (fake *Checker) TextMarshalerArgsForCall(i int) types.Type { + fake.textMarshalerMutex.RLock() + defer fake.textMarshalerMutex.RUnlock() + argsForCall := fake.textMarshalerArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *Checker) TextMarshalerReturns(result1 bool) { + fake.textMarshalerMutex.Lock() + defer fake.textMarshalerMutex.Unlock() + fake.TextMarshalerStub = nil + fake.textMarshalerReturns = struct { + result1 bool + }{result1} +} + +func (fake *Checker) TextMarshalerReturnsOnCall(i int, result1 bool) { + fake.textMarshalerMutex.Lock() + defer fake.textMarshalerMutex.Unlock() + fake.TextMarshalerStub = nil + if fake.textMarshalerReturnsOnCall == nil { + fake.textMarshalerReturnsOnCall = make(map[int]struct { + result1 bool + }) + } + fake.textMarshalerReturnsOnCall[i] = struct { + result1 bool + }{result1} +} + +func (fake *Checker) TextUnmarshaler(arg1 types.Type) bool { + fake.textUnmarshalerMutex.Lock() + ret, specificReturn := fake.textUnmarshalerReturnsOnCall[len(fake.textUnmarshalerArgsForCall)] + fake.textUnmarshalerArgsForCall = append(fake.textUnmarshalerArgsForCall, struct { + arg1 types.Type + }{arg1}) + stub := fake.TextUnmarshalerStub + fakeReturns := fake.textUnmarshalerReturns + fake.recordInvocation("TextUnmarshaler", []interface{}{arg1}) + fake.textUnmarshalerMutex.Unlock() + if stub != nil { + return stub(arg1) + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *Checker) TextUnmarshalerCallCount() int { + fake.textUnmarshalerMutex.RLock() + defer fake.textUnmarshalerMutex.RUnlock() + return len(fake.textUnmarshalerArgsForCall) +} + +func (fake *Checker) TextUnmarshalerCalls(stub func(types.Type) bool) { + fake.textUnmarshalerMutex.Lock() + defer fake.textUnmarshalerMutex.Unlock() + fake.TextUnmarshalerStub = stub +} + +func (fake *Checker) TextUnmarshalerArgsForCall(i int) types.Type { + fake.textUnmarshalerMutex.RLock() + defer fake.textUnmarshalerMutex.RUnlock() + argsForCall := fake.textUnmarshalerArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *Checker) TextUnmarshalerReturns(result1 bool) { + fake.textUnmarshalerMutex.Lock() + defer fake.textUnmarshalerMutex.Unlock() + fake.TextUnmarshalerStub = nil + fake.textUnmarshalerReturns = struct { + result1 bool + }{result1} +} + +func (fake *Checker) TextUnmarshalerReturnsOnCall(i int, result1 bool) { + fake.textUnmarshalerMutex.Lock() + defer fake.textUnmarshalerMutex.Unlock() + fake.TextUnmarshalerStub = nil + if fake.textUnmarshalerReturnsOnCall == nil { + fake.textUnmarshalerReturnsOnCall = make(map[int]struct { + result1 bool + }) + } + fake.textUnmarshalerReturnsOnCall[i] = struct { + result1 bool + }{result1} +} + +func (fake *Checker) Invocations() map[string][][]interface{} { + fake.invocationsMutex.RLock() + defer fake.invocationsMutex.RUnlock() + copiedInvocations := map[string][][]interface{}{} + for key, value := range fake.invocations { + copiedInvocations[key] = value + } + return copiedInvocations +} + +func (fake *Checker) recordInvocation(key string, args []interface{}) { + fake.invocationsMutex.Lock() + defer fake.invocationsMutex.Unlock() + if fake.invocations == nil { + fake.invocations = map[string][][]interface{}{} + } + if fake.invocations[key] == nil { + fake.invocations[key] = [][]interface{}{} + } + fake.invocations[key] = append(fake.invocations[key], args) +} + +var _ muxt.Checker = new(Checker) diff --git a/internal/generate/file.go b/internal/generate/file.go index 134dff69..6abe6edc 100644 --- a/internal/generate/file.go +++ b/internal/generate/file.go @@ -6,7 +6,6 @@ import ( "go/ast" "go/parser" "go/token" - "go/types" "log" "maps" "path" @@ -40,17 +39,17 @@ func newFile(pkg source.Package) *File { // 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) - return parser.ParseExpr(s) +// TypeExpr spells t as the generated file refers to it, importing what it +// needs. +func (file *File) TypeExpr(t source.Type) (ast.Expr, error) { + return parser.ParseExpr(t.Format(file.qualify)) } -// pkgQualifier implements types.Qualifier -func (file *File) pkgQualifier(pkg *types.Package) string { - if pkg.Path() == file.pkg.Types.Path() { +func (file *File) qualify(pkgName, pkgPath string) string { + if pkgPath == file.pkg.Types.Path() { return "" } - return file.Import(pkg.Name(), pkg.Path()) + return file.Import(pkgName, pkgPath) } func (file *File) Import(pkgIdent, pkgPath string) string { @@ -62,7 +61,7 @@ func (file *File) Import(pkgIdent, pkgPath string) string { } func (file *File) ImportSpecs() []*ast.ImportSpec { - result := append(make([]*ast.ImportSpec, 0, len(file.importSpecs)), file.importSpecs...) + result := slices.Clone(file.importSpecs) slices.SortFunc(result, func(a, b *ast.ImportSpec) int { return strings.Compare(a.Path.Value, b.Path.Value) }) return slices.CompactFunc(result, func(a, b *ast.ImportSpec) bool { return a.Path.Value == b.Path.Value }) } 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 fbb5676f..f07db0e8 100644 --- a/internal/generate/groups.go +++ b/internal/generate/groups.go @@ -6,7 +6,6 @@ import ( "strings" "github.com/typelate/muxt/internal/muxt" - "github.com/typelate/muxt/internal/source" ) type templateGroups struct { @@ -15,32 +14,24 @@ type templateGroups struct { all []muxt.Definition } -func groupTemplates(config RoutesFileConfiguration, variables []source.Variable) (templateGroups, error) { +func groupTemplates(config RoutesFileConfiguration, defs []muxt.Definition) (templateGroups, error) { result := templateGroups{ byFile: make(map[string][]muxt.Definition), + all: defs, } - for _, tv := range variables { - defs, err := muxt.Definitions(tv) - if err != nil { - return result, err - } - - if !config.OutputDatastar { - for _, d := range defs { - if d.UsesSignals() { - return result, fmt.Errorf("the signals argument in %q requires --output-datastar; it is shorthand for unmarshalJSON(body)", d.Name()) - } - if name, ok := d.SignalsCallback(); ok { - return result, fmt.Errorf("the %s callback in %q requires --output-datastar; it marshals its argument as a datastar-patch-signals event", name, d.Name()) - } - } - } - + if !config.OutputDatastar { for _, d := range defs { - key := d.SourceFile() - result.byFile[key] = append(result.byFile[key], d) + if d.UsesSignals() { + return result, fmt.Errorf("the signals argument in %q requires --output-datastar; it is shorthand for unmarshalJSON(body)", d.Name()) + } + if name, ok := d.SignalsCallback(); ok { + return result, fmt.Errorf("the %s callback in %q requires --output-datastar; it marshals its argument as a datastar-patch-signals event", name, d.Name()) + } } - result.all = append(result.all, defs...) + } + for _, d := range defs { + key := d.SourceFile() + result.byFile[key] = append(result.byFile[key], d) } if err := muxt.CheckForDuplicatePatterns(result.all); err != nil { diff --git a/internal/generate/html.go b/internal/generate/html.go index ae22faf5..a22534dd 100644 --- a/internal/generate/html.go +++ b/internal/generate/html.go @@ -3,54 +3,145 @@ package generate import ( "go/ast" "go/token" - "go/types" "net/http" "slices" "strconv" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) -// executeHTMLTemplateHandler assembles a rendered-route handler: template +// newHTMLTemplateHandler assembles a rendered-route handler: template // data, argument parsing, the method call, and template execution into the // response buffer. The optional respond statements run after execution and // before the status/body write, letting another representation (marshalJSON) // replace the buffered output on success without duplicating the assembly. -func executeHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def muxt.Definition, sig *types.Signature, resultDataIdent string, receiverInterfaceName string, bufIdent string, statusCodeIdent string, respond ...ast.Stmt) (*ast.FuncLit, error) { - var callFun ast.Expr - isMethodCall := sig.Recv() != nil - if isMethodCall { - callFun = &ast.SelectorExpr{ - X: ast.NewIdent(receiverIdent), - Sel: ast.NewIdent(def.FunctionIdentifier().Name), - } - } else { - callFun = ast.NewIdent(def.FunctionIdentifier().Name) - } - - execIdx, hasExecute := -1, false - var resultType types.Type - var execHasArg bool - for i, arg := range def.Arguments { - if arg.Type == muxt.ArgumentTypeExecute && arg.Identifier == muxt.TemplateNameScopeIdentifierExecute { - // The callback contract (func() error or func(T) error) is - // validated by muxt.ResolveCall, which records T and whether the - // callback takes the data argument. - execIdx, hasExecute = i, true - resultType, execHasArg = arg.CallbackResultType(), arg.CallbackHasArg() - break - } +func newHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def muxt.Definition, resultDataIdent string, receiverInterfaceName string, bufIdent string, statusCodeIdent string, respond ...ast.Stmt) (*ast.FuncLit, error) { + if execIdx, ok := def.ExecuteArgumentIndex(); ok { + return newExecuteHTMLTemplateHandler(file, config, def, execIdx, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent, respond...) + } + return newResultHTMLTemplateHandler(file, config, def, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent, respond...) +} + +func newExecuteHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def muxt.Definition, execIdx int, resultDataIdent string, receiverInterfaceName string, bufIdent string, statusCodeIdent string, respond ...ast.Stmt) (*ast.FuncLit, error) { + callFun := callFuncExpression(def) + resultType := def.ResultType() + execHasArg := def.Arguments[execIdx].CallbackHasArg() + + handlerFunc, call, lit, err := initHandlerScope(file, config, def, resultDataIdent, receiverInterfaceName, bufIdent, resultType) + if err != nil { + return lit, err + } + + const guardIdent = "executed" + closure, err := executeClosure(file, def, resultDataIdent, bufIdent, guardIdent, resultType, execHasArg) + if err != nil { + return nil, err + } + // The render callback may be invoked more than once (possibly from + // another goroutine); guard with an atomic.Bool so it renders at most + // once (see executeClosure). + handlerFunc.Body.List = append(handlerFunc.Body.List, &ast.DeclStmt{Decl: &ast.GenDecl{ + Tok: token.VAR, + Specs: []ast.Spec{&ast.ValueSpec{Names: []*ast.Ident{ast.NewIdent(guardIdent)}, Type: astgen.ExportedIdentifier(file, "", "sync/atomic", "Bool")}}, + }}) + callArgs := slices.Clone(call.Args) + callArgs[execIdx] = closure + if config.Logger { + handlerFunc.Body.List = append(handlerFunc.Body.List, logDebugStatement(file, "handling request", def.RawPattern())) + } + renderCheck := checkExecuteTemplateError(file, config.Logger, def.RawPattern()) + renderCheck.Init = &ast.AssignStmt{ + Lhs: []ast.Expr{ast.NewIdent(errIdent)}, + Tok: token.DEFINE, + Rhs: []ast.Expr{&ast.CallExpr{Fun: callFun, Args: callArgs}}, } - if !hasExecute { - resultType = sig.Results().At(0).Type() + setOkay := &ast.AssignStmt{ + Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(resultDataIdent), Sel: ast.NewIdent(TemplateDataFieldIdentifierOkay)}}, + Tok: token.ASSIGN, + Rhs: []ast.Expr{astgen.Bool(true)}, + } + handlerFunc.Body.List = append(handlerFunc.Body.List, &ast.IfStmt{ + Cond: &ast.BinaryExpr{ + X: astgen.CallBuiltinLen(&ast.SelectorExpr{X: ast.NewIdent(resultDataIdent), Sel: ast.NewIdent(TemplateDataFieldIdentifierError)}), + Op: token.EQL, + Y: astgen.Int(0), + }, + Body: &ast.BlockStmt{List: []ast.Stmt{renderCheck, setOkay}}, + }) + + handlerFunc.Body.List = append(handlerFunc.Body.List, respond...) + + return writeHeadersAndStatusCode(file, handlerFunc, def, statusCodeIdent, bufIdent, resultDataIdent) +} + +func newResultHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def muxt.Definition, resultDataIdent string, receiverInterfaceName string, bufIdent string, statusCodeIdent string, respond ...ast.Stmt) (*ast.FuncLit, error) { + callFun := callFuncExpression(def) + resultType := def.ResultType() + handlerFunc, call, lit, err := initHandlerScope(file, config, def, resultDataIdent, receiverInterfaceName, bufIdent, resultType) + if err != nil { + return lit, err } - typeExpr, err := file.TypeASTExpression(resultType) + + errBody := appendTemplateDataError(file, resultDataIdent, ast.NewIdent(errIdent)) + errBody.List = append(errBody.List, assignTemplateDataErrStatusCode(file, resultDataIdent, http.StatusInternalServerError)) + receiverCall, err := callReceiverMethod(resultDataIdent, &ast.SelectorExpr{ + X: ast.NewIdent(resultDataIdent), + Sel: ast.NewIdent(TemplateDataFieldIdentifierResult), + }, def.ResultShape(), def.FunctionIdentifier().Name, &ast.CallExpr{ + Fun: callFun, + Args: slices.Clone(call.Args), + }, errBody) if err != nil { return nil, err } + handlerFunc.Body.List = append(handlerFunc.Body.List, &ast.IfStmt{ + Cond: &ast.BinaryExpr{ + X: astgen.CallBuiltinLen(&ast.SelectorExpr{ + X: ast.NewIdent(resultDataIdent), + Sel: ast.NewIdent(TemplateDataFieldIdentifierError), + }), + Op: token.EQL, + Y: astgen.Int(0), + }, + Body: &ast.BlockStmt{ + List: receiverCall.Stmts(), + }, + }) + + callExecuteTemplate(file, config, def, handlerFunc, bufIdent, resultDataIdent) + + handlerFunc.Body.List = append(handlerFunc.Body.List, respond...) + + return writeHeadersAndStatusCode(file, handlerFunc, def, statusCodeIdent, bufIdent, resultDataIdent) +} + +func initHandlerScope(file *File, config RoutesFileConfiguration, def muxt.Definition, resultDataIdent string, receiverInterfaceName string, bufIdent string, resultType source.Type) (*ast.FuncLit, *ast.CallExpr, *ast.FuncLit, error) { + typeExpr, err := file.TypeExpr(resultType) + if err != nil { + return nil, nil, nil, err + } + + handlerFunc := newHandlerFuncLit(file, config, resultDataIdent, receiverInterfaceName, typeExpr) - handlerFunc := &ast.FuncLit{ + // Parsing rewrites the call's arguments to the locals it declares, + // so it works on a copy and the definition stays as resolved. + call := def.CallExpression() + if handlerFunc.Body.List, err = appendParseArgumentStatements(handlerFunc.Body.List, def, file, 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 + }, nil); err != nil { + return nil, nil, nil, err + } + + handlerFunc.Body.List = append(handlerFunc.Body.List, astgen.GetBufferFromPool(file, bufferPoolIdent, bufIdent)...) + return handlerFunc, call, nil, nil +} + +func newHandlerFuncLit(file *File, config RoutesFileConfiguration, resultDataIdent, receiverInterfaceName string, typeExpr ast.Expr) *ast.FuncLit { + return &ast.FuncLit{ Type: astgen.HTTPHandlerFuncType(file, muxt.TemplateNameScopeIdentifierHTTPResponse, muxt.TemplateNameScopeIdentifierHTTPRequest), Body: &ast.BlockStmt{ List: []ast.Stmt{ @@ -74,91 +165,11 @@ func executeHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def }, }, } +} - // 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 - }, nil); err != nil { - return nil, err - } - - handlerFunc.Body.List = append(handlerFunc.Body.List, astgen.GetBufferFromPool(file, bufferPoolIdent, bufIdent)...) - - if hasExecute { - const guardIdent = "executed" - closure, err := executeClosure(file, def, resultDataIdent, bufIdent, guardIdent, resultType, execHasArg) - if err != nil { - return nil, err - } - // The render callback may be invoked more than once (possibly from - // another goroutine); guard with an atomic.Bool so it renders at most - // once (see executeClosure). - handlerFunc.Body.List = append(handlerFunc.Body.List, &ast.DeclStmt{Decl: &ast.GenDecl{ - Tok: token.VAR, - Specs: []ast.Spec{&ast.ValueSpec{Names: []*ast.Ident{ast.NewIdent(guardIdent)}, Type: astgen.ExportedIdentifier(file, "", "sync/atomic", "Bool")}}, - }}) - callArgs := slices.Clone(call.Args) - callArgs[execIdx] = closure - if config.Logger { - handlerFunc.Body.List = append(handlerFunc.Body.List, logDebugStatement(file, "handling request", def.RawPattern())) - } - renderCheck := checkExecuteTemplateError(file, config.Logger, def.RawPattern()) - renderCheck.Init = &ast.AssignStmt{ - Lhs: []ast.Expr{ast.NewIdent(errIdent)}, - Tok: token.DEFINE, - Rhs: []ast.Expr{&ast.CallExpr{Fun: callFun, Args: callArgs}}, - } - setOkay := &ast.AssignStmt{ - Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(resultDataIdent), Sel: ast.NewIdent(TemplateDataFieldIdentifierOkay)}}, - Tok: token.ASSIGN, - Rhs: []ast.Expr{astgen.Bool(true)}, - } - handlerFunc.Body.List = append(handlerFunc.Body.List, &ast.IfStmt{ - Cond: &ast.BinaryExpr{ - X: astgen.CallBuiltinLen(&ast.SelectorExpr{X: ast.NewIdent(resultDataIdent), Sel: ast.NewIdent(TemplateDataFieldIdentifierError)}), - Op: token.EQL, - Y: astgen.Int(0), - }, - Body: &ast.BlockStmt{List: []ast.Stmt{renderCheck, setOkay}}, - }) - } else { - errBody := appendTemplateDataError(file, resultDataIdent, ast.NewIdent(errIdent)) - errBody.List = append(errBody.List, assignTemplateDataErrStatusCode(file, resultDataIdent, http.StatusInternalServerError)) - receiverCall, err := callReceiverMethod(resultDataIdent, &ast.SelectorExpr{ - X: ast.NewIdent(resultDataIdent), - Sel: ast.NewIdent(TemplateDataFieldIdentifierResult), - }, sig, def.FunctionIdentifier().Name, &ast.CallExpr{ - Fun: callFun, - Args: slices.Clone(call.Args), - }, errBody) - if err != nil { - return nil, err - } - handlerFunc.Body.List = append(handlerFunc.Body.List, &ast.IfStmt{ - Cond: &ast.BinaryExpr{ - X: astgen.CallBuiltinLen(&ast.SelectorExpr{ - X: ast.NewIdent(resultDataIdent), - Sel: ast.NewIdent(TemplateDataFieldIdentifierError), - }), - Op: token.EQL, - Y: astgen.Int(0), - }, - Body: &ast.BlockStmt{ - List: receiverCall.Stmts(), - }, - }) - - callExecuteTemplate(file, config, def, handlerFunc, bufIdent, resultDataIdent) - } - - handlerFunc.Body.List = append(handlerFunc.Body.List, respond...) - +func writeHeadersAndStatusCode(file *File, handlerFunc *ast.FuncLit, def muxt.Definition, statusCodeIdent string, bufIdent string, resultDataIdent string) (*ast.FuncLit, error) { if !def.HasResponseWriterArg() { - handlerFunc.Body.List = append(handlerFunc.Body.List, writeStatusAndHeaders(file, def, resultType, def.DefaultStatusCode(), statusCodeIdent, bufIdent, resultDataIdent, func() ast.Expr { + handlerFunc.Body.List = append(handlerFunc.Body.List, writeStatusAndHeaders(file, def, def.DefaultStatusCode(), statusCodeIdent, bufIdent, resultDataIdent, func() ast.Expr { return &ast.SelectorExpr{X: ast.NewIdent(resultDataIdent), Sel: ast.NewIdent(TemplateDataFieldIdentifierResult)} })...) } else { @@ -167,6 +178,16 @@ func executeHTMLTemplateHandler(file *File, config RoutesFileConfiguration, def return handlerFunc, nil } +func callFuncExpression(def muxt.Definition) ast.Expr { + if !def.IsMethod() { + return ast.NewIdent(def.FunctionIdentifier().Name) + } + return &ast.SelectorExpr{ + X: ast.NewIdent(receiverIdent), + Sel: ast.NewIdent(def.FunctionIdentifier().Name), + } +} + func callExecuteTemplate(file *File, config RoutesFileConfiguration, def muxt.Definition, handlerFunc *ast.FuncLit, bufIdent string, dataIdent string) { if config.Logger { handlerFunc.Body.List = append(handlerFunc.Body.List, logDebugStatement(file, "handling request", def.RawPattern())) @@ -203,7 +224,7 @@ func callExecuteTemplate(file *File, config RoutesFileConfiguration, def muxt.De // than once gets an error on the later calls rather than a second render. The // guard is an atomic.Bool compared-and-swapped so a callback invoked from // another goroutine still renders exactly once. -func executeClosure(file *File, def muxt.Definition, tdIdent, bufIdent, guardIdent string, resultType types.Type, hasArg bool) (*ast.FuncLit, error) { +func executeClosure(file *File, def muxt.Definition, tdIdent, bufIdent, guardIdent string, resultType source.Type, hasArg bool) (*ast.FuncLit, error) { const dataIdent = "data" var params []*ast.Field body := []ast.Stmt{ @@ -218,7 +239,7 @@ func executeClosure(file *File, def muxt.Definition, tdIdent, bufIdent, guardIde }, } if hasArg { - tExpr, err := file.TypeASTExpression(resultType) + tExpr, err := file.TypeExpr(resultType) if err != nil { return nil, err } diff --git a/internal/generate/marshal_json.go b/internal/generate/marshal_json.go index 4e40ecc1..83d4f4c1 100644 --- a/internal/generate/marshal_json.go +++ b/internal/generate/marshal_json.go @@ -4,8 +4,8 @@ import ( "fmt" "go/ast" "go/token" - "go/types" "net/http" + "slices" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/muxt" @@ -18,13 +18,13 @@ import ( // the method succeeded the rendered output is discarded and the marshaled // result is written as application/json; on any recorded error the rendered // output is sent as the usual text/html fallback. -func marshalJSONHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.Definition, sig *types.Signature, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent string) (*ast.FuncLit, error) { - for _, arg := range def.Arguments { - if arg.Type == muxt.ArgumentTypeExecute && arg.Identifier == muxt.TemplateNameScopeIdentifierExecute { - return nil, fmt.Errorf("marshalJSON does not support the execute callback") - } +func marshalJSONHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.Definition, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent string) (*ast.FuncLit, error) { + if slices.ContainsFunc(def.Arguments, func(arg muxt.Argument) bool { + return arg.Type == muxt.ArgumentTypeExecute && arg.Identifier == muxt.TemplateNameScopeIdentifierExecute + }) { + return nil, fmt.Errorf("marshalJSON does not support the execute callback") } - return executeHTMLTemplateHandler(file, config, def, sig, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent, marshalJSONRespondStmts(file, resultDataIdent, bufIdent)...) + return newHTMLTemplateHandler(file, config, def, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent, marshalJSONRespondStmts(file, resultDataIdent, bufIdent)...) } // marshalJSONRespondStmts builds: diff --git a/internal/generate/routes.go b/internal/generate/routes.go index 5653dba9..5a49f8b7 100644 --- a/internal/generate/routes.go +++ b/internal/generate/routes.go @@ -5,7 +5,6 @@ import ( "fmt" "go/ast" "go/token" - "go/types" "log" "maps" "net/http" @@ -16,7 +15,6 @@ import ( "github.com/ettle/strcase" - "github.com/typelate/muxt/internal/asteval" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/muxt" "github.com/typelate/muxt/internal/source" @@ -99,9 +97,8 @@ const DefaultMultipartMaxMemory int64 = 32 << 20 // 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) { +// directory. defs are pkg's route definitions, resolved by muxt.ResolveDefinitions. +func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, pkg source.Package, defs []muxt.Definition, logger *log.Logger) ([]GeneratedFile, error) { if !token.IsIdentifier(config.PackageName) { return nil, fmt.Errorf("package name %q is not an identifier", config.PackageName) } @@ -112,20 +109,15 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, pkg source.P config.PackageName = pkg.Types.Name() config.SSETemplateDataType = cmp.Or(config.SSETemplateDataType, "SSETemplateData") - if receiver == nil { - receiver = asteval.NamedEmptyStruct("Receiver", pkg.Types) - } - - groups, err := groupTemplates(config, pkg.Variables) + groups, err := groupTemplates(config, defs) if err != nil { return nil, err } var ( receiverInterface = &ast.InterfaceType{Methods: new(ast.FieldList)} - templateSourceFiles = slices.Collect(maps.Keys(groups.byFile)) + templateSourceFiles = slices.Sorted(maps.Keys(groups.byFile)) ) - slices.Sort(templateSourceFiles) // Build main routes function routesFunc := &ast.FuncDecl{ @@ -178,7 +170,7 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, pkg source.P generatedFiles []GeneratedFile ) if config.OutputMultipleFiles { - files, err := sourceFileRouteFunctionFiles(wd, config, templateSourceFiles, groups, logger, file, receiver, receiverInterface, routesFunc) + files, err := sourceFileRouteFunctionFiles(wd, config, templateSourceFiles, groups, logger, file, receiverInterface, routesFunc) if err != nil { return files, err } @@ -192,8 +184,8 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, pkg source.P routesFunc.Body.List = append(routesFunc.Body.List, bytesBufferPoolDeclaration(file)) } - // Generate handlers for parse-based templates (empty sourceFile) - if err := hydrateGroup(topLevelTemplateRoutes, file, receiver, receiverInterface, logger, config.ReceiverType != "" && logger != nil, logger != nil && !config.SilenceHTTPResponseWarning); err != nil { + logResolutionNotes(topLevelTemplateRoutes, config, logger) + if err := collectReceiverMethods(topLevelTemplateRoutes, file, receiverInterface); err != nil { return nil, err } for _, def := range topLevelTemplateRoutes { @@ -282,30 +274,24 @@ func TemplateRoutesFiles(wd string, config RoutesFileConfiguration, pkg source.P // 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, receiverInterface *ast.InterfaceType, logger *log.Logger, noteSynthesized, warnResponse bool) error { - var resolveErrs []error +// logResolutionNotes reports what resolution found for defs. +func logResolutionNotes(defs []muxt.Definition, config RoutesFileConfiguration, logger *log.Logger) { + if logger == nil { + return + } synthesized := 0 - for i := range defs { - if warnResponse && defs[i].HasResponseWriterArg() { + for _, def := range defs { + if !config.SilenceHTTPResponseWarning && def.HasResponseWriterArg() { // Taking over the http.ResponseWriter is an escape hatch: // muxt then leaves the response entirely to the method. - logger.Printf("warning: %s 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", defs[i].Pattern()) + logger.Printf("warning: %s 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", def.Pattern()) } - if defs[i].FunctionIdentifier() == nil { + if config.ReceiverType == "" { continue } - if err := muxt.ResolveCall(&defs[i], file.OutputPackage(), receiver); err != nil { - resolveErrs = append(resolveErrs, err) - continue - } - if noteSynthesized { - for _, sig := range defs[i].SynthesizedMethods() { - logger.Printf("note: %s does not define %s", receiver.Obj().Name(), sig) - synthesized++ - } - } - if err := accumulateReceiverMethods(defs[i].FunctionIdentifier().Name, defs[i].Signature(), defs[i].IsMethod(), defs[i].Arguments, file, receiverInterface); err != nil { - return err + for _, sig := range def.SynthesizedMethods() { + logger.Printf("note: %s does not define %s", config.ReceiverType, sig) + synthesized++ } } if synthesized > 0 { @@ -313,10 +299,23 @@ func hydrateGroup(defs []muxt.Definition, file *File, receiver *types.Named, rec // on .Result are deferred; say so once. logger.Printf("note: the inferred signatures return any — implement the methods to type-check the templates against real types") } - return muxt.CombineErrors(resolveErrs) } -func accumulateReceiverMethods(name string, sig *types.Signature, isMethod bool, args []muxt.Argument, file *File, receiverInterface *ast.InterfaceType) error { +// collectReceiverMethods adds the receiver methods the calls in defs need to +// receiverInterface. +func collectReceiverMethods(defs []muxt.Definition, file *File, receiverInterface *ast.InterfaceType) error { + for _, def := range defs { + if def.FunctionIdentifier() == nil { + continue + } + if err := accumulateReceiverMethods(def.FunctionIdentifier().Name, def.Signature(), def.IsMethod(), def.Arguments, file, receiverInterface); err != nil { + return err + } + } + return nil +} + +func accumulateReceiverMethods(name string, sig source.Type, isMethod bool, args []muxt.Argument, file *File, receiverInterface *ast.InterfaceType) error { // Recurse into nested call arguments regardless of whether this call is a // receiver method: a package-scope function may receive nested receiver // method calls that must appear in the interface. @@ -331,12 +330,12 @@ func accumulateReceiverMethods(name string, sig *types.Signature, isMethod bool, if !isMethod { return nil } - if i := slices.IndexFunc(receiverInterface.Methods.List, func(field *ast.Field) bool { + if slices.ContainsFunc(receiverInterface.Methods.List, func(field *ast.Field) bool { return field.Names[0].Name == name - }); i >= 0 { + }) { return nil } - exp, err := file.TypeASTExpression(sig) + exp, err := file.TypeExpr(sig) if err != nil { return err } @@ -347,7 +346,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, receiverInterface *ast.InterfaceType, routesFunc *ast.FuncDecl) ([]GeneratedFile, error) { +func sourceFileRouteFunctionFiles(wd string, config RoutesFileConfiguration, templateSourceFiles []string, groups templateGroups, logger *log.Logger, file *File, receiverInterface *ast.InterfaceType, routesFunc *ast.FuncDecl) ([]GeneratedFile, error) { var generatedFiles []GeneratedFile for _, sourceFile := range templateSourceFiles { definitions := groups.byFile[sourceFile] @@ -359,7 +358,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, newFile(file.OutputPackage()), routesFuncName, receiverInterfaceName, logger, config, receiver) + perFileAST, err := generatePerFileAST(sourceFile, definitions, newFile(file.OutputPackage()), routesFuncName, receiverInterfaceName, logger, config) if err != nil { return nil, fmt.Errorf("failed to generate routes for %s: %w", sourceFile, err) } @@ -479,7 +478,6 @@ func generatePerFileRouteFunction( receiverInterfaceName string, logger *log.Logger, config RoutesFileConfiguration, - receiver *types.Named, receiverInterface *ast.InterfaceType, ) (*ast.FuncDecl, error) { if sourceFile == "" { @@ -525,7 +523,8 @@ func generatePerFileRouteFunction( } // Generate handlers for each template - if err := hydrateGroup(defs, file, receiver, receiverInterface, logger, config.ReceiverType != "" && logger != nil, logger != nil && !config.SilenceHTTPResponseWarning); err != nil { + logResolutionNotes(defs, config, logger) + if err := collectReceiverMethods(defs, file, receiverInterface); err != nil { return nil, err } for i := range defs { @@ -559,7 +558,6 @@ func generatePerFileAST( funcName, receiverInterfaceName string, logger *log.Logger, config RoutesFileConfiguration, - receiver *types.Named, ) (*ast.File, error) { if sourceFile == "" { return nil, fmt.Errorf("sourceFile cannot be empty") @@ -578,7 +576,6 @@ func generatePerFileAST( receiverInterfaceName, logger, config, - receiver, scopedReceiverInterface, ) if err != nil { @@ -654,7 +651,7 @@ func noReceiverMethodCall(file *File, def muxt.Definition, config RoutesFileConf callExecuteTemplate(file, config, def, handlerFunc, bufIdent, templateDataVarIdent) - handlerFunc.Body.List = append(handlerFunc.Body.List, writeStatusAndHeaders(file, def, types.NewStruct(nil, nil), def.DefaultStatusCode(), statusCodeIdent, bufIdent, templateDataVarIdent, func() ast.Expr { + handlerFunc.Body.List = append(handlerFunc.Body.List, writeStatusAndHeaders(file, def, def.DefaultStatusCode(), statusCodeIdent, bufIdent, templateDataVarIdent, func() ast.Expr { panic("when no receiver method is called, then the result variable should not be needed") })...) return handlerFunc @@ -666,17 +663,17 @@ func callHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.Defini statusCodeIdent = "statusCode" resultDataIdent = "td" ) - sig := def.Signature() - if sig == nil { + + if def.Signature().IsZero() { return nil, fmt.Errorf("call for pattern %s was not resolved", def.Pattern()) } switch def.Representation { case muxt.RepresentationSSE: - return sseMethodHandlerFunc(file, config, def, sig, receiverInterfaceName) + return sseMethodHandlerFunc(file, config, def, receiverInterfaceName) case muxt.RepresentationMarshalJSON: - return marshalJSONHandlerFunc(file, config, def, sig, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent) + return marshalJSONHandlerFunc(file, config, def, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent) default: - return executeHTMLTemplateHandler(file, config, def, sig, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent) + return newHTMLTemplateHandler(file, config, def, resultDataIdent, receiverInterfaceName, bufIdent, statusCodeIdent) } } @@ -747,28 +744,7 @@ 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) { +func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, file *File, 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 // data (and respond with the recorded error status). SSE handlers pass @@ -779,21 +755,14 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f if !ok { return nil, fmt.Errorf("expected function to be identifier") } - if signature == nil { - return nil, fmt.Errorf("call %s was not resolved to a signature", fun.Name) - } - // const parsedVariableName = "parsed" - if exp := signature.Params().Len(); exp != len(call.Args) { // TODO: (signature.Variadic() && exp > len(call.Args)) - sigStr := fun.Name + strings.TrimPrefix(signature.String(), "func") - return nil, fmt.Errorf("handler func %s expects %d arguments but call %s has %d", sigStr, signature.Params().Len(), astgen.Format(call), len(call.Args)) + if len(args) != len(call.Args) { + return nil, fmt.Errorf("call %s was not resolved", fun.Name) } if parsed == nil { parsed = make(map[string]struct{}) } resultCount := 0 for i, a := range call.Args { - param := signature.Params().At(i) - switch arg := a.(type) { default: // TODO: add error case @@ -801,7 +770,7 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f nestedArg := args[i] if nestedArg.Type == muxt.ArgumentTypeRequestBodyJSON { const bodyValueIdent = "bodyValue" - decodeStatements, err := decodeJSONBodyStatements(file, bodyValueIdent, nestedArg.ParamType, parseErrBlock) + decodeStatements, err := decodeJSONBodyStatements(file, bodyValueIdent, nestedArg.ParamType(), parseErrBlock) if err != nil { return nil, err } @@ -809,7 +778,7 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f call.Args[i] = ast.NewIdent(bodyValueIdent) continue } - parseArgStatements, err := appendParseArgumentStatements(statements, def, file, resultType, nestedArg.Signature(), nestedArg.Arguments(), parsed, rdIdent, config, arg, validationFailureBlock, parseErrBlock) + parseArgStatements, err := appendParseArgumentStatements(statements, def, file, nestedArg.Arguments(), parsed, rdIdent, config, arg, validationFailureBlock, parseErrBlock) if err != nil { return nil, err } @@ -819,11 +788,6 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f funcIdent := arg.Fun.(*ast.Ident).Name - callSig := nestedArg.Signature() - if callSig == nil { - return nil, fmt.Errorf("call %s was not resolved to a signature", funcIdent) - } - if nestedArg.IsMethod() { arg.Fun = &ast.SelectorExpr{ X: ast.NewIdent(receiverIdent), @@ -835,7 +799,7 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f errBody := appendTemplateDataError(file, rdIdent, ast.NewIdent(errIdent)) errBody.List = append(errBody.List, assignTemplateDataErrStatusCode(file, rdIdent, http.StatusInternalServerError)) - nestedCall, err := callReceiverMethod(rdIdent, ast.NewIdent(resultVarIdent), callSig, funcIdent, arg, errBody) + nestedCall, err := callReceiverMethod(rdIdent, ast.NewIdent(resultVarIdent), nestedArg.ResultShape(), funcIdent, arg, errBody) if err != nil { return nil, err } @@ -849,28 +813,25 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f continue } name := arg.Name - argType, ok := muxt.DefaultScopeType(file.OutputPackage(), &def, name) - if !ok { - return nil, fmt.Errorf("failed to determine type for %s", name) - } + argument := args[i] src := requestArgumentSource(def, name) ident := name - if slices.Contains(def.PathValueIdentifiers(), name) { + if def.ArgumentIsPathParameter(name) { ident = pathParamIdent(name) call.Args[i] = ast.NewIdent(ident) } - if types.AssignableTo(argType, param.Type()) { + if argument.Direct() { if _, ok := parsed[name]; !ok { parsed[name] = struct{}{} switch name { case muxt.TemplateNameScopeIdentifierForm: - declareFormVar, err := formVariableAssignment(file, arg, param.Type()) + declareFormVar, err := formVariableAssignment(file, arg, argument.ParamType()) if err != nil { return nil, err } statements = append(statements, callParseForm(file), declareFormVar) case muxt.TemplateNameScopeIdentifierMultipart: - declareMultipartVar, err := multipartVariableAssignment(file, arg, param.Type()) + declareMultipartVar, err := multipartVariableAssignment(file, arg, argument.ParamType()) if err != nil { return nil, err } @@ -880,7 +841,7 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f case muxt.TemplateNameScopeIdentifierRequestBody: statements = append(statements, singleAssignment(token.DEFINE, ast.NewIdent(ident))(src)) default: - if slices.Contains(def.PathValueIdentifiers(), name) || name == muxt.TemplateNameScopeIdentifierLastEventID { + if def.ArgumentIsPathParameter(name) || name == muxt.TemplateNameScopeIdentifierLastEventID { statements = append(statements, singleAssignment(token.DEFINE, ast.NewIdent(ident))(src)) } } @@ -891,35 +852,38 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f continue } switch { - case slices.Contains(def.PathValueIdentifiers(), name): + case def.ArgumentIsPathParameter(name): parsed[name] = struct{}{} - s, err := generateParseValueFromStringStatements(file, def, name+"Parsed", resultType, src, param.Type(), nil, singleAssignment(token.DEFINE, ast.NewIdent(ident)), parseErrBlock()) + s, err := generateParseValueFromStringStatements(file, name+"Parsed", src, argument.ParamType(), argument.UnmarshalMethod(), nil, singleAssignment(token.DEFINE, ast.NewIdent(ident)), parseErrBlock()) if err != nil { return nil, err } statements = append(statements, s...) 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()) + s, err := generateParseValueFromStringStatements(file, name+"Parsed", src, argument.ParamType(), argument.UnmarshalMethod(), nil, singleAssignment(token.DEFINE, ast.NewIdent(ident)), parseErrBlock()) if err != nil { return nil, err } statements = append(statements, s...) case arg.Name == muxt.TemplateNameScopeIdentifierForm: - s, err := appendParseFormToStructStatements(statements, def, file, resultType, arg, args[i], validationFailureBlock, parseErrBlock) + s, err := appendParseFormToStructStatements(statements, file, arg, argument, validationFailureBlock, parseErrBlock) if err != nil { return nil, err } statements = s case arg.Name == muxt.TemplateNameScopeIdentifierMultipart: - s, err := appendParseMultipartFormToStructStatements(statements, def, file, resultType, arg, args[i], validationFailureBlock, parseErrBlock, config) + s, err := appendParseMultipartFormToStructStatements(statements, file, arg, argument, validationFailureBlock, parseErrBlock, config) if err != nil { return nil, err } statements = s default: - pt, _ := file.TypeASTExpression(param.Type()) - at, _ := file.TypeASTExpression(argType) + if argument.ScopeType().IsZero() { + return nil, fmt.Errorf("failed to determine type for %s", name) + } + pt, _ := file.TypeExpr(argument.ParamType()) + at, _ := file.TypeExpr(argument.ScopeType()) return nil, fmt.Errorf("method expects type %s but %s is %s", astgen.Format(pt), arg.Name, astgen.Format(at)) } } @@ -927,19 +891,19 @@ func appendParseArgumentStatements(statements []ast.Stmt, def muxt.Definition, f return statements, nil } -func appendParseFormToStructStatements(statements []ast.Stmt, def muxt.Definition, file *File, resultType types.Type, arg *ast.Ident, argument muxt.Argument, validationBlock ValidationErrorBlock, parseErrBlock func() *ast.BlockStmt) ([]ast.Stmt, error) { - return appendStructFieldParseStatements(statements, def, file, resultType, arg, argument, validationBlock, parseErrBlock, callParseForm(file)) +func appendParseFormToStructStatements(statements []ast.Stmt, file *File, arg *ast.Ident, argument muxt.Argument, validationBlock ValidationErrorBlock, parseErrBlock func() *ast.BlockStmt) ([]ast.Stmt, error) { + return appendStructFieldParseStatements(statements, file, arg, argument, validationBlock, parseErrBlock, callParseForm(file)) } // appendStructFieldParseStatements renders the per-field parse statements for // a form or multipart struct parameter from the field bindings resolved by // muxt.ResolveCall. Used by both `form` (parseCall = callParseForm(file)) and // `multipart` (parseCall = callParseMultipartForm(...)). -func appendStructFieldParseStatements(statements []ast.Stmt, def muxt.Definition, file *File, resultType types.Type, arg *ast.Ident, argument muxt.Argument, validationBlock ValidationErrorBlock, parseErrBlock func() *ast.BlockStmt, parseCall ast.Stmt) ([]ast.Stmt, error) { +func appendStructFieldParseStatements(statements []ast.Stmt, file *File, arg *ast.Ident, argument muxt.Argument, validationBlock ValidationErrorBlock, parseErrBlock func() *ast.BlockStmt, parseCall ast.Stmt) ([]ast.Stmt, error) { const parsedVariableName = "value" statements = append(statements, parseCall) - declareVar, err := formVariableDeclaration(file, arg, argument.ParamType) + declareVar, err := formVariableDeclaration(file, arg, argument.ParamType()) if err != nil { return nil, err } @@ -948,9 +912,9 @@ func appendStructFieldParseStatements(statements []ast.Stmt, def muxt.Definition for _, fb := range argument.FormFields() { if fb.FileHeader { if fb.Slice { - statements = append(statements, fileHeaderSliceAssignment(arg, fb.Field.Name(), fb.InputName)) + statements = append(statements, fileHeaderSliceAssignment(arg, fb.Name, fb.InputName)) } else { - statements = append(statements, fileHeaderSingleAssignment(arg, fb.Field.Name(), fb.InputName)) + statements = append(statements, fileHeaderSingleAssignment(arg, fb.Name, fb.InputName)) } continue } @@ -959,14 +923,14 @@ func appendStructFieldParseStatements(statements []ast.Stmt, def muxt.Definition if fb.Slice { parseResult := func(expr ast.Expr) ast.Stmt { return &ast.AssignStmt{ - Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(arg.Name), Sel: ast.NewIdent(fb.Field.Name())}}, + Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(arg.Name), Sel: ast.NewIdent(fb.Name)}}, Tok: token.ASSIGN, - Rhs: []ast.Expr{astgen.CallBuiltinAppend(&ast.SelectorExpr{X: ast.NewIdent(arg.Name), Sel: ast.NewIdent(fb.Field.Name())}, expr)}, + Rhs: []ast.Expr{astgen.CallBuiltinAppend(&ast.SelectorExpr{X: ast.NewIdent(arg.Name), Sel: ast.NewIdent(fb.Name)}, expr)}, } } - parseStatements, err := generateParseValueFromStringStatements(file, def, parsedVariableName, resultType, ast.NewIdent("val"), fb.Elem, validations, parseResult, parseErrBlock()) + parseStatements, err := generateParseValueFromStringStatements(file, parsedVariableName, ast.NewIdent("val"), fb.Elem(), fb.Method, validations, parseResult, parseErrBlock()) if err != nil { - return nil, fmt.Errorf("failed to generate parse statements for %s field %s: %w", arg.Name, fb.Field.Name(), err) + return nil, fmt.Errorf("failed to generate parse statements for %s field %s: %w", arg.Name, fb.Name, err) } statements = append(statements, &ast.RangeStmt{ Key: ast.NewIdent("_"), @@ -978,15 +942,15 @@ func appendStructFieldParseStatements(statements []ast.Stmt, def muxt.Definition } else { parseResult := func(expr ast.Expr) ast.Stmt { return &ast.AssignStmt{ - Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(arg.Name), Sel: ast.NewIdent(fb.Field.Name())}}, + Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(arg.Name), Sel: ast.NewIdent(fb.Name)}}, Tok: token.ASSIGN, Rhs: []ast.Expr{expr}, } } str := &ast.CallExpr{Fun: &ast.SelectorExpr{X: ast.NewIdent(muxt.TemplateNameScopeIdentifierHTTPRequest), Sel: ast.NewIdent("FormValue")}, Args: []ast.Expr{&ast.BasicLit{Kind: token.STRING, Value: strconv.Quote(fb.InputName)}}} - parseStatements, err := generateParseValueFromStringStatements(file, def, parsedVariableName, resultType, str, fb.Elem, validations, parseResult, parseErrBlock()) + parseStatements, err := generateParseValueFromStringStatements(file, parsedVariableName, str, fb.Elem(), fb.Method, validations, parseResult, parseErrBlock()) if err != nil { - return nil, fmt.Errorf("failed to generate parse statements for %s field %s: %w", arg.Name, fb.Field.Name(), err) + return nil, fmt.Errorf("failed to generate parse statements for %s field %s: %w", arg.Name, fb.Name, err) } if len(parseStatements) > 1 { statements = append(statements, &ast.BlockStmt{ @@ -1006,8 +970,8 @@ func appendStructFieldParseStatements(statements []ast.Stmt, def muxt.Definition // FileHeader field bindings (from request.MultipartForm.File) are resolved by // muxt.ResolveCall; all other field-binding behavior is shared with the form // codepath. -func appendParseMultipartFormToStructStatements(statements []ast.Stmt, def muxt.Definition, file *File, resultType types.Type, arg *ast.Ident, argument muxt.Argument, validationBlock ValidationErrorBlock, parseErrBlock func() *ast.BlockStmt, config RoutesFileConfiguration) ([]ast.Stmt, error) { - return appendStructFieldParseStatements(statements, def, file, resultType, arg, argument, validationBlock, parseErrBlock, callParseMultipartForm(file, config, parseErrBlock())) +func appendParseMultipartFormToStructStatements(statements []ast.Stmt, file *File, arg *ast.Ident, argument muxt.Argument, validationBlock ValidationErrorBlock, parseErrBlock func() *ast.BlockStmt, config RoutesFileConfiguration) ([]ast.Stmt, error) { + return appendStructFieldParseStatements(statements, file, arg, argument, validationBlock, parseErrBlock, callParseMultipartForm(file, config, parseErrBlock())) } // fileHeaderSingleAssignment emits: @@ -1081,8 +1045,8 @@ func wrapInMultipartFormNotNil(stmt ast.Stmt) ast.Stmt { } } -func formVariableDeclaration(file *File, arg *ast.Ident, tp types.Type) (*ast.DeclStmt, error) { - typeExp, err := file.TypeASTExpression(tp) +func formVariableDeclaration(file *File, arg *ast.Ident, tp source.Type) (*ast.DeclStmt, error) { + typeExp, err := file.TypeExpr(tp) if err != nil { return nil, err } @@ -1099,8 +1063,8 @@ func formVariableDeclaration(file *File, arg *ast.Ident, tp types.Type) (*ast.De }, nil } -func formVariableAssignment(file *File, arg *ast.Ident, tp types.Type) (*ast.DeclStmt, error) { - typeExp, err := file.TypeASTExpression(tp) +func formVariableAssignment(file *File, arg *ast.Ident, tp source.Type) (*ast.DeclStmt, error) { + typeExp, err := file.TypeExpr(tp) if err != nil { return nil, err } @@ -1144,16 +1108,20 @@ func templateDataParseErrBlock(file *File, rdIdent string) *ast.BlockStmt { // errBlock, which callers supply so the failure can be handled differently per // context (normal handlers accumulate into the template data; SSE handlers // respond 400 before establishing the stream). -func generateParseValueFromStringStatements(file *File, _ muxt.Definition, tmp string, _ types.Type, str ast.Expr, valueType types.Type, validations []ast.Stmt, assignment func(ast.Expr) ast.Stmt, errBlock *ast.BlockStmt) ([]ast.Stmt, error) { +func generateParseValueFromStringStatements(file *File, tmp string, str ast.Expr, valueType source.Type, method muxt.UnmarshalMethod, validations []ast.Stmt, assignment func(ast.Expr) ast.Stmt, errBlock *ast.BlockStmt) ([]ast.Stmt, error) { + typeExpr, err := file.TypeExpr(valueType) + if err != nil { + return nil, err + } // convert wraps the parsed value in a conversion to the target basic type // for the strconv functions that return a wider type (ParseInt/ParseUint). convert := func(exp ast.Expr) ast.Stmt { return assignment(&ast.CallExpr{ - Fun: ast.NewIdent(valueType.(*types.Basic).Name()), + Fun: typeExpr, Args: []ast.Expr{exp}, }) } - switch muxt.UnmarshalMethodFor(file.OutputPackage(), valueType) { + switch method { case muxt.UnmarshalBool: return parseBlock(tmp, astgen.StrconvParseBoolCall(file, str), validations, errBlock, assignment), nil case muxt.UnmarshalInt: @@ -1193,7 +1161,6 @@ func generateParseValueFromStringStatements(file *File, _ muxt.Definition, tmp s }}, validations, []ast.Stmt{assignment(ast.NewIdent(tmp))}) return statements, nil case muxt.UnmarshalTextUnmarshaler: - tp, _ := file.TypeASTExpression(valueType) return []ast.Stmt{ &ast.DeclStmt{ Decl: &ast.GenDecl{ @@ -1201,7 +1168,7 @@ func generateParseValueFromStringStatements(file *File, _ muxt.Definition, tmp s Specs: []ast.Spec{ &ast.ValueSpec{ Names: []*ast.Ident{ast.NewIdent(tmp)}, - Type: tp, + Type: typeExpr, }, }, }, @@ -1233,8 +1200,7 @@ func generateParseValueFromStringStatements(file *File, _ muxt.Definition, tmp s assignment(ast.NewIdent(tmp)), }, nil default: - tp, _ := file.TypeASTExpression(valueType) - return nil, fmt.Errorf("unsupported type: %s", astgen.Format(tp)) + return nil, fmt.Errorf("unsupported type: %s", astgen.Format(typeExpr)) } } @@ -1291,14 +1257,12 @@ func (r *receiverMethodCall) Stmts() []ast.Stmt { return stmts } -func callReceiverMethod(rdIdent string, dataVar ast.Expr, method *types.Signature, callIdent string, call *ast.CallExpr, errBody *ast.BlockStmt) (*receiverMethodCall, error) { - const ( - okIdent = "ok" - ) - switch method.Results().Len() { +func callReceiverMethod(rdIdent string, dataVar ast.Expr, shape muxt.ResultShape, callIdent string, call *ast.CallExpr, errBody *ast.BlockStmt) (*receiverMethodCall, error) { + const okIdent = "ok" + switch shape { default: return nil, fmt.Errorf("method %s has no results it should have one or two", callIdent) - case 1: + case muxt.ResultShapeData: return &receiverMethodCall{ Assign: &ast.AssignStmt{Lhs: []ast.Expr{dataVar}, Tok: token.ASSIGN, Rhs: []ast.Expr{call}}, SetOkay: &ast.AssignStmt{Lhs: []ast.Expr{&ast.SelectorExpr{ @@ -1306,42 +1270,28 @@ func callReceiverMethod(rdIdent string, dataVar ast.Expr, method *types.Signatur Sel: ast.NewIdent(TemplateDataFieldIdentifierOkay), }}, Tok: token.ASSIGN, Rhs: []ast.Expr{astgen.Bool(true)}}, }, nil - case 2: - lastResult := method.Results().At(method.Results().Len() - 1).Type() - - errorType := types.Universe.Lookup("error").Type().Underlying().(*types.Interface) - - if types.Implements(lastResult, errorType) { - return &receiverMethodCall{ - VarDecl: &ast.DeclStmt{Decl: &ast.GenDecl{Tok: token.VAR, Specs: []ast.Spec{&ast.ValueSpec{Names: []*ast.Ident{ast.NewIdent(errIdent)}, Type: ast.NewIdent("error")}}}}, - Assign: &ast.AssignStmt{Lhs: []ast.Expr{dataVar, ast.NewIdent(errIdent)}, Tok: token.ASSIGN, Rhs: []ast.Expr{call}}, - Check: &ast.IfStmt{ - Cond: &ast.BinaryExpr{X: ast.NewIdent(errIdent), Op: token.NEQ, Y: astgen.Nil()}, - Body: errBody, - }, - }, nil - } - - if basic, ok := lastResult.(*types.Basic); ok && basic.Kind() == types.Bool { - return &receiverMethodCall{ - VarDecl: &ast.DeclStmt{Decl: &ast.GenDecl{Tok: token.VAR, Specs: []ast.Spec{&ast.ValueSpec{Names: []*ast.Ident{ast.NewIdent("ok")}, Type: ast.NewIdent("bool")}}}}, - Assign: &ast.AssignStmt{Lhs: []ast.Expr{dataVar, ast.NewIdent(okIdent)}, Tok: token.ASSIGN, Rhs: []ast.Expr{call}}, - Check: &ast.IfStmt{ - Cond: &ast.UnaryExpr{Op: token.NOT, X: ast.NewIdent(okIdent)}, - Body: &ast.BlockStmt{ - List: []ast.Stmt{ - &ast.ReturnStmt{}, - }, - }, - }, - SetOkay: &ast.AssignStmt{Lhs: []ast.Expr{&ast.SelectorExpr{ - X: ast.NewIdent(rdIdent), - Sel: ast.NewIdent(TemplateDataFieldIdentifierOkay), - }}, Tok: token.ASSIGN, Rhs: []ast.Expr{astgen.Bool(true)}}, - }, nil - } - - return nil, fmt.Errorf("expected last result to be either an error or a bool") + case muxt.ResultShapeDataError: + return &receiverMethodCall{ + VarDecl: &ast.DeclStmt{Decl: &ast.GenDecl{Tok: token.VAR, Specs: []ast.Spec{&ast.ValueSpec{Names: []*ast.Ident{ast.NewIdent(errIdent)}, Type: ast.NewIdent("error")}}}}, + Assign: &ast.AssignStmt{Lhs: []ast.Expr{dataVar, ast.NewIdent(errIdent)}, Tok: token.ASSIGN, Rhs: []ast.Expr{call}}, + Check: &ast.IfStmt{ + Cond: &ast.BinaryExpr{X: ast.NewIdent(errIdent), Op: token.NEQ, Y: astgen.Nil()}, + Body: errBody, + }, + }, nil + case muxt.ResultShapeDataOK: + return &receiverMethodCall{ + VarDecl: &ast.DeclStmt{Decl: &ast.GenDecl{Tok: token.VAR, Specs: []ast.Spec{&ast.ValueSpec{Names: []*ast.Ident{ast.NewIdent(okIdent)}, Type: ast.NewIdent("bool")}}}}, + Assign: &ast.AssignStmt{Lhs: []ast.Expr{dataVar, ast.NewIdent(okIdent)}, Tok: token.ASSIGN, Rhs: []ast.Expr{call}}, + Check: &ast.IfStmt{ + Cond: &ast.UnaryExpr{Op: token.NOT, X: ast.NewIdent(okIdent)}, + Body: &ast.BlockStmt{List: []ast.Stmt{&ast.ReturnStmt{}}}, + }, + SetOkay: &ast.AssignStmt{Lhs: []ast.Expr{&ast.SelectorExpr{ + X: ast.NewIdent(rdIdent), + Sel: ast.NewIdent(TemplateDataFieldIdentifierOkay), + }}, Tok: token.ASSIGN, Rhs: []ast.Expr{astgen.Bool(true)}}, + }, nil } } @@ -1350,8 +1300,8 @@ func callReceiverMethod(rdIdent string, dataVar ast.Expr, method *types.Signatur // // var bodyValue T // if err := json.NewDecoder(request.Body).Decode(&bodyValue); err != nil { } -func decodeJSONBodyStatements(file *File, valueIdent string, paramType types.Type, parseErrBlock func() *ast.BlockStmt) ([]ast.Stmt, error) { - typeExpr, err := file.TypeASTExpression(paramType) +func decodeJSONBodyStatements(file *File, valueIdent string, paramType source.Type, parseErrBlock func() *ast.BlockStmt) ([]ast.Stmt, error) { + typeExpr, err := file.TypeExpr(paramType) if err != nil { return nil, err } @@ -1395,7 +1345,7 @@ func requestArgumentSource(def muxt.Definition, name string) ast.Expr { Sel: ast.NewIdent("Body"), } } - if name == muxt.TemplateNameScopeIdentifierLastEventID && !slices.Contains(def.PathValueIdentifiers(), name) { + if def.ArgumentIsLastEventID(name) { return &ast.CallExpr{ Fun: &ast.SelectorExpr{ X: &ast.SelectorExpr{X: ast.NewIdent(muxt.TemplateNameScopeIdentifierHTTPRequest), Sel: ast.NewIdent("Header")}, @@ -1489,8 +1439,8 @@ func callParseMultipartForm(file *File, config RoutesFileConfiguration, errBlock // multipartVariableAssignment emits `var = request.MultipartForm` // for raw-mode multipart binding. -func multipartVariableAssignment(file *File, arg *ast.Ident, tp types.Type) (*ast.DeclStmt, error) { - typeExp, err := file.TypeASTExpression(tp) +func multipartVariableAssignment(file *File, arg *ast.Ident, tp source.Type) (*ast.DeclStmt, error) { + typeExp, err := file.TypeExpr(tp) if err != nil { return nil, err } @@ -1536,16 +1486,15 @@ func singleAssignment(assignTok token.Token, result ast.Expr) func(exp ast.Expr) } } -var statusCoder = statusCoderInterface() - -func writeStatusAndHeaders(file *File, def muxt.Definition, resultType types.Type, fallbackStatusCode int, statusCode, bufIdent, resultDataIdent string, resultVar func() ast.Expr) []ast.Stmt { +func writeStatusAndHeaders(file *File, def muxt.Definition, fallbackStatusCode int, statusCode, bufIdent, resultDataIdent string, resultVar func() ast.Expr) []ast.Stmt { statusCodePriorityList := []ast.Expr{ &ast.SelectorExpr{X: ast.NewIdent(resultDataIdent), Sel: ast.NewIdent(templateDataFieldStatusCode)}, &ast.SelectorExpr{X: ast.NewIdent(resultDataIdent), Sel: ast.NewIdent(TemplateDataFieldIdentifierErrStatusCode)}, } - if types.Implements(resultType, statusCoder) { + switch def.ResultStatusCode() { + case muxt.ResultStatusCodeMethod: statusCodePriorityList = append(statusCodePriorityList, &ast.CallExpr{Fun: &ast.SelectorExpr{X: resultVar(), Sel: ast.NewIdent("StatusCode")}}) - } else if obj, _, _ := types.LookupFieldOrMethod(resultType, true, file.OutputPackage().Types, "StatusCode"); obj != nil { + case muxt.ResultStatusCodeField: statusCodePriorityList = append(statusCodePriorityList, &ast.SelectorExpr{X: resultVar(), Sel: ast.NewIdent("StatusCode")}) } var list []ast.Stmt @@ -1666,16 +1615,6 @@ func logDebugStatement(file *File, message, pattern string) *ast.ExprStmt { } } -func statusCoderInterface() *types.Interface { - sig := types.NewSignatureType(nil, nil, nil, - types.NewTuple(), - types.NewTuple(types.NewVar(token.NoPos, nil, "", types.Typ[types.Int])), - false) - - method := types.NewFunc(token.NoPos, nil, "StatusCode", sig) - return types.NewInterfaceType([]*types.Func{method}, nil).Complete() -} - func assignTemplateDataErrStatusCode(file *File, rdIdent string, code int) *ast.AssignStmt { return &ast.AssignStmt{ Lhs: []ast.Expr{&ast.SelectorExpr{ diff --git a/internal/generate/routes_test.go b/internal/generate/routes_test.go new file mode 100644 index 00000000..9104bfc7 --- /dev/null +++ b/internal/generate/routes_test.go @@ -0,0 +1,78 @@ +package generate + +import ( + "io" + "log" + "strings" + "testing" + + "github.com/typelate/muxt/internal/astgen" + "github.com/typelate/muxt/internal/fake" + "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}}`) + + defs, err := muxt.ResolveDefinitions(pkg, receiver, fake.NewChecker().Fake()) + if err != nil { + t.Fatal(err) + } + groups, err := groupTemplates(config, defs) + if err != nil { + t.Fatal(err) + } + file := newFile(pkg) + def := groups.all[0] + + 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) + } +} + +// TestHandlerGenerationRejectsAnArgumentWithNoRequestValue generates an sse +// handler whose call passes a message, which names no request value to +// parse. +func TestHandlerGenerationRejectsAnArgumentWithNoRequestValue(t *testing.T) { + config := testConfig() + config.ReceiverType = "T" + pkg, receiver := testSource(t, `package server + +type T struct{} + +func (T) Stream(string) {} +`, "T", `{{define "GET /x sse(Stream(fooMessage))"}}{{end}}{{define "fooMessage"}}{{end}}`) + + defs, err := muxt.ResolveDefinitions(pkg, receiver, fake.NewChecker().Fake()) + if err != nil { + t.Fatal(err) + } + _, err = TemplateRoutesFiles(".", config, pkg, defs, log.New(io.Discard, "", 0)) + if err == nil || !strings.Contains(err.Error(), "failed to determine type for fooMessage") { + t.Errorf("got error %v, want it to say it failed to determine type for fooMessage", err) + } +} diff --git a/internal/generate/snapshot_test.go b/internal/generate/snapshot_test.go new file mode 100644 index 00000000..978843a2 --- /dev/null +++ b/internal/generate/snapshot_test.go @@ -0,0 +1,250 @@ +package generate_test + +import ( + "encoding/json/v2" + "errors" + "flag" + "go/ast" + "go/parser" + "go/token" + "log" + "maps" + "os" + "path" + "path/filepath" + "slices" + "strconv" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "golang.org/x/tools/txtar" + + "github.com/typelate/muxt/internal/configjson" + "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/generate and compares them with the archive's want/ files. The +// directory is the command, so one case runs with +// -run TestSnapshots/generate/sse. +// +// An archive holds everything a case is: the configuration it generates +// with, its inputs, and what it generates. +// +// - config.json is the configuration, as the command line named in the +// archive's header parses into it. TestCommandLineConfigurations in +// internal/cli states that parse; this states what generate does with +// the result. +// - Go files and .gohtml files are loaded as example.com/server by +// internal/load/loadtest -- type checked against the official standard +// library, without loading the package graph -- 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) { + const command = "generate" + if stray, _ := filepath.Glob(filepath.Join("testdata", "*.txtar")); len(stray) > 0 { + t.Fatalf("%s is not in a command's directory; generate's archives are in testdata/%s", stray[0], command) + } + archives, err := filepath.Glob(filepath.Join("testdata", command, "*.txtar")) + if err != nil { + t.Fatal(err) + } + if len(archives) == 0 { + t.Fatalf("no archives in testdata/%s", command) + } + t.Run(command, func(t *testing.T) { + for _, archivePath := range archives { + runSnapshot(t, archivePath) + } + }) +} + +// runSnapshot compares one archive's want/ files with what it generates, +// or with -update rewrites them. +func runSnapshot(t *testing.T, archivePath string) { + t.Helper() + t.Run(strings.TrimSuffix(filepath.Base(archivePath), ".txtar"), func(t *testing.T) { + archive, err := txtar.ParseFile(archivePath) + if err != nil { + t.Fatal(err) + } + got := snapshot(t, configuration(t, archive), 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) { + assert.Equal(t, want[name], got[name], "want/%s differs (run go test -run TestSnapshots -update to rewrite)") + } + }) +} + +// configuration reads the archive's config.json: the configuration to +// generate with. +func configuration(t *testing.T, archive *txtar.Archive) generate.RoutesFileConfiguration { + t.Helper() + for _, file := range archive.Files { + if file.Name != "config.json" { + continue + } + var config generate.RoutesFileConfiguration + if err := json.Unmarshal(file.Data, &config, configjson.Options()); err != nil { + t.Fatalf("config.json: %v", err) + } + return config + } + t.Fatal("the archive has no config.json") + return generate.RoutesFileConfiguration{} +} + +// 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/") || file.Name == "config.json" { + 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 := errors.AsType[muxt.MultiLineError](err); ok { + text = multiLine.MultiLineError() + } + got["error.txt"] = relative(text) + "\n" + return got + } + + pl := loadtest.Package(t, dir, "example.com/server", files) + pkg, receiver, err := load.GenerateSource(dir, pl, config) + if err != nil { + return fail(err) + } + defs, err := muxt.ResolveDefinitions(pkg, receiver, load.StandardLibrary(pl)) + if err != nil { + return fail(err) + } + var logs strings.Builder + generated, err := generate.TemplateRoutesFiles(dir, config, pkg, defs, 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(ms ...map[string]string) []string { + var keys []string + for _, m := range ms { + keys = slices.AppendSeq(keys, maps.Keys(m)) + } + slices.Sort(keys) + return slices.Compact(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/source_test.go b/internal/generate/source_test.go new file mode 100644 index 00000000..02de0067 --- /dev/null +++ b/internal/generate/source_test.go @@ -0,0 +1,52 @@ +package generate + +import ( + "go/types" + "html/template" + "testing" + + "github.com/typelate/muxt/internal/fake" + "github.com/typelate/muxt/internal/source" +) + +// This file builds what generation reads without loading a module: a +// package type checked in memory, with no 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. The source imports nothing; resolving routes over it takes a +// a muxttest checker. +func testSource(t *testing.T, goSource, receiverType, templates string) (source.Package, *types.Named) { + t.Helper() + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": goSource}) + src := source.Package{ + Fset: fake.FileSet, + Types: pkg, + Variables: []source.Variable{{ + Name: "templates", + Set: template.Must(template.New("templates").Parse(templates)), + }}, + } + if receiverType == "" { + return src, nil + } + return src, fake.Lookup(t, pkg, receiverType).(*types.Named) +} diff --git a/internal/generate/sse.go b/internal/generate/sse.go index 96a29fd9..0caf4c54 100644 --- a/internal/generate/sse.go +++ b/internal/generate/sse.go @@ -3,20 +3,20 @@ package generate import ( "go/ast" "go/token" - "go/types" "net/http" "slices" "strconv" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) // sseMethodHandlerFunc builds the http.HandlerFunc for a route that streams // Server-Sent Events. Unlike a normal handler it establishes an event stream // (Content-Type text/event-stream, flush) and invokes the receiver method with // a callback closure that renders and writes one SSE frame per call. -func sseMethodHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.Definition, sig *types.Signature, receiverInterfaceName string) (*ast.FuncLit, error) { +func sseMethodHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.Definition, receiverInterfaceName string) (*ast.FuncLit, error) { const ( flusherIdent = "flusher" okIdent = "ok" @@ -79,12 +79,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. // 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) + call := def.CallExpression() + body, err := appendParseArgumentStatements(body, def, file, def.Arguments, nil, "", config, call, validationFailureBlock, parseErrBlock) if err != nil { return nil, err } @@ -136,6 +134,7 @@ func sseMethodHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.D return nil, err } callArgs[i] = closure + default: } } callExpr := &ast.CallExpr{Fun: callFun, Args: callArgs} @@ -170,7 +169,7 @@ func sseMethodHandlerFunc(file *File, config RoutesFileConfiguration, def muxt.D // flusher.Flush() // return nil // } -func signalsClosure(file *File, resultType types.Type, flusherIdent, mutexIdent string) (*ast.FuncLit, error) { +func signalsClosure(file *File, resultType source.Type, flusherIdent, mutexIdent string) (*ast.FuncLit, error) { const ( resultIdent = "result" payloadIdent = "payload" @@ -178,7 +177,7 @@ func signalsClosure(file *File, resultType types.Type, flusherIdent, mutexIdent response := muxt.TemplateNameScopeIdentifierHTTPResponse request := muxt.TemplateNameScopeIdentifierHTTPRequest - resultTypeExpr, err := file.TypeASTExpression(resultType) + resultTypeExpr, err := file.TypeExpr(resultType) if err != nil { return nil, err } @@ -256,7 +255,7 @@ func requestContextCancelledCheck(request string) ast.Stmt { // } // // For the zero-arg form it omits the parameter and the result field. -func sseClosure(file *File, config RoutesFileConfiguration, def muxt.Definition, templateName string, resultType types.Type, hasArg bool, receiverInterfaceName, flusherIdent, mutexIdent string) (*ast.FuncLit, error) { +func sseClosure(file *File, config RoutesFileConfiguration, def muxt.Definition, templateName string, resultType source.Type, hasArg bool, receiverInterfaceName, flusherIdent, mutexIdent string) (*ast.FuncLit, error) { const ( bufIdent = "buf" tdIdent = "td" @@ -265,7 +264,7 @@ func sseClosure(file *File, config RoutesFileConfiguration, def muxt.Definition, response := muxt.TemplateNameScopeIdentifierHTTPResponse request := muxt.TemplateNameScopeIdentifierHTTPRequest - resultTypeExpr, err := file.TypeASTExpression(resultType) + resultTypeExpr, err := file.TypeExpr(resultType) if err != nil { return nil, err } diff --git a/internal/generate/template_route_path.go b/internal/generate/template_route_path.go index f02d8be7..fd8fb016 100644 --- a/internal/generate/template_route_path.go +++ b/internal/generate/template_route_path.go @@ -6,12 +6,11 @@ import ( "fmt" "go/ast" "go/token" - "go/types" "strconv" - "strings" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" ) const ( @@ -59,16 +58,6 @@ 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.OutputPackage().Import("encoding") - if !ok { - return nil, false, false, fmt.Errorf(`the "encoding" package must be loaded`) - } - scope := encodingPkg.Scope() - textMarshalerObject := scope.Lookup("TextMarshaler") - textMarshalerType := textMarshalerObject.Type() - textMarshalerUnderlying := textMarshalerType.Underlying() - textMarshalerInterface := textMarshalerUnderlying.(*types.Interface) - ident, err := def.ExportedPathIdentifier() if err != nil { return nil, false, false, err @@ -90,34 +79,13 @@ func routePathFunc(file *File, config RoutesFileConfiguration, def *muxt.Definit }, } - if def.Path() == "/" || def.Path() == "/{$}" { - if config.PathPrefix { - method.Body.List = []ast.Stmt{&ast.ReturnStmt{Results: []ast.Expr{ - astgen.Call(file, "path", "path", "Join", - astgen.Call(file, "cmp", "cmp", "Or", - &ast.SelectorExpr{ - X: ast.NewIdent(methodReceiverName), - Sel: ast.NewIdent(pathPrefixPathsStructFieldName), - }, - astgen.String("/"), - ), - ), - }}} - } else { - method.Body.List = []ast.Stmt{&ast.ReturnStmt{Results: []ast.Expr{astgen.String("/")}}} - } - return method, usesEscaper, usesSegmentsEscaper, nil + if def.IsIndex() { + return indexRoutePath(file, config, method, methodReceiverName, usesEscaper, usesSegmentsEscaper) } - templatePath, hasDollarSuffix := strings.CutSuffix(def.Path(), "{$}") - segmentStrings := strings.Split(templatePath, "/") var ( fields []*ast.Field - last types.Type - - identIndex = 0 - - segmentIdentifiers = def.PathValueIdentifiers() + last source.Type ) hasErrorResult := false @@ -130,39 +98,35 @@ func routePathFunc(file *File, config RoutesFileConfiguration, def *muxt.Definit astgen.String("/"), ), } - for si, segment := range segmentStrings { - if len(segment) < 1 { - continue - } - if segment[0] != '{' || segment[len(segment)-1] != '}' { + for i, segment := range def.Segments { + // si numbers the segment as the path splits on "/", with the empty + // segment before the leading "/" at 0. + si := i + 1 + if segment.IsLiteral() { if len(segmentExpressions) > 0 { prev := segmentExpressions[len(segmentExpressions)-1] if prevBasic, ok := prev.(*ast.BasicLit); ok { prevVal, _ := strconv.Unquote(prevBasic.Value) - prevBasic.Value = strconv.Quote(prevVal + "/" + segment) + prevBasic.Value = strconv.Quote(prevVal + "/" + segment.Value()) continue } } segmentExpressions = append(segmentExpressions, &ast.BasicLit{ Kind: token.STRING, - Value: strconv.Quote(segment), + Value: strconv.Quote(segment.Value()), }) continue } - name := segmentIdentifiers[identIndex] + name := segment.Value() ident := pathParamIdent(name) - wildcard := si == len(segmentStrings)-1 && isWildcardSegment(segment) - pathValueType, ok := def.ArgumentType(name) - identIndex++ - if !ok { - pathValueType = types.Universe.Lookup("string").Type() - } - tpNode, err := file.TypeASTExpression(pathValueType) + wildcard := segment.IsRemainder() + pathValueType := segment.Type() + tpNode, err := file.TypeExpr(pathValueType) if err != nil { return nil, false, false, err } - if last != nil && len(fields) > 0 && types.Identical(last, pathValueType) { + if len(fields) > 0 && last.Identical(pathValueType) { fields[len(fields)-1].Names = append(fields[len(fields)-1].Names, ast.NewIdent(ident)) } else { fields = append(fields, &ast.Field{ @@ -176,7 +140,7 @@ func routePathFunc(file *File, config RoutesFileConfiguration, def *muxt.Definit summer.Write([]byte(def.Name())) pathHash := hex.EncodeToString(summer.Sum(nil)) - if types.Implements(pathValueType, textMarshalerInterface) { + if segment.TextMarshaler() { hasErrorResult = true if len(method.Type.Results.List) == 1 { method.Type.Results.List = append(method.Type.Results.List, &ast.Field{ @@ -227,15 +191,11 @@ func routePathFunc(file *File, config RoutesFileConfiguration, def *muxt.Definit continue } - basicType, ok := pathValueType.Underlying().(*types.Basic) - if !ok { - return nil, false, false, fmt.Errorf("unsupported type %s for path parameters: %s", astgen.Format(tpNode), ident) - } - exp, err := astgen.ConvertToString(file, ast.NewIdent(ident), basicType.Kind()) + exp, err := astgen.ConvertToString(file, ast.NewIdent(ident), pathValueType) if err != nil { return nil, false, false, fmt.Errorf("failed to encode variable %s: %v", ident, err) } - if basicType.Info()&types.IsString != 0 { + if pathValueType.IsString() { if wildcard { usesSegmentsEscaper = true exp = escapedPathSegments(exp) @@ -254,7 +214,7 @@ func routePathFunc(file *File, config RoutesFileConfiguration, def *muxt.Definit }, Args: segmentExpressions, }) - if hasDollarSuffix { + if def.HasPathEndWildcard() { returnStmt = &ast.BinaryExpr{ X: returnStmt, Op: token.ADD, @@ -276,11 +236,23 @@ func routePathFunc(file *File, config RoutesFileConfiguration, def *muxt.Definit return method, usesEscaper, usesSegmentsEscaper, nil } -// isWildcardSegment reports whether segment is a {name...} pattern; only the -// trailing segment of a pattern may be one, and its value names a path suffix -// spliced without escaping. -func isWildcardSegment(segment string) bool { - return strings.HasSuffix(strings.TrimSuffix(segment, "}"), "...") +func indexRoutePath(file *File, config RoutesFileConfiguration, method *ast.FuncDecl, methodReceiverName string, usesEscaper bool, usesSegmentsEscaper bool) (*ast.FuncDecl, bool, bool, error) { + if config.PathPrefix { + method.Body.List = []ast.Stmt{&ast.ReturnStmt{Results: []ast.Expr{ + astgen.Call(file, "path", "path", "Join", + astgen.Call(file, "cmp", "cmp", "Or", + &ast.SelectorExpr{ + X: ast.NewIdent(methodReceiverName), + Sel: ast.NewIdent(pathPrefixPathsStructFieldName), + }, + astgen.String("/"), + ), + ), + }}} + } else { + method.Body.List = []ast.Stmt{&ast.ReturnStmt{Results: []ast.Expr{astgen.String("/")}}} + } + return method, usesEscaper, usesSegmentsEscaper, nil } // pathParamIdent names the generated local for a path parameter. The suffix diff --git a/internal/generate/testdata/generate/err_duplicate_pattern.txtar b/internal/generate/testdata/generate/err_duplicate_pattern.txtar new file mode 100644 index 00000000..6ea98c9a --- /dev/null +++ b/internal/generate/testdata/generate/err_duplicate_pattern.txtar @@ -0,0 +1,40 @@ +Two templates may not register the same pattern. + +Command line: muxt generate +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/err_name_errors.txtar b/internal/generate/testdata/generate/err_name_errors.txtar new file mode 100644 index 00000000..741d0d17 --- /dev/null +++ b/internal/generate/testdata/generate/err_name_errors.txtar @@ -0,0 +1,49 @@ +Every malformed name is reported, each at its position. + +Command line: muxt generate +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/err_resolution.txtar b/internal/generate/testdata/generate/err_resolution.txtar new file mode 100644 index 00000000..89bd24c0 --- /dev/null +++ b/internal/generate/testdata/generate/err_resolution.txtar @@ -0,0 +1,61 @@ +Resolution errors from every route are reported together, pointing at +the argument or method at fault, with where the method is defined. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/err_response_state_with_response_argument.txtar b/internal/generate/testdata/generate/err_response_state_with_response_argument.txtar new file mode 100644 index 00000000..a5dd0b36 --- /dev/null +++ b/internal/generate/testdata/generate/err_response_state_with_response_argument.txtar @@ -0,0 +1,44 @@ +A template may not set the status of a route whose method took the +response. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/err_route_paths_method_collision.txtar b/internal/generate/testdata/generate/err_route_paths_method_collision.txtar new file mode 100644 index 00000000..1f7bcf8e --- /dev/null +++ b/internal/generate/testdata/generate/err_route_paths_method_collision.txtar @@ -0,0 +1,41 @@ +Handlers whose names export to the same TemplateRoutePaths method +collide. + +Command line: muxt generate +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/err_signals_without_datastar.txtar b/internal/generate/testdata/generate/err_signals_without_datastar.txtar new file mode 100644 index 00000000..1623d484 --- /dev/null +++ b/internal/generate/testdata/generate/err_signals_without_datastar.txtar @@ -0,0 +1,37 @@ +signals is Datastar's request body, so it needs --output-datastar. + +Command line: muxt generate +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/execute_callback.txtar b/internal/generate/testdata/generate/execute_callback.txtar new file mode 100644 index 00000000..0e384bc9 --- /dev/null +++ b/internal/generate/testdata/generate/execute_callback.txtar @@ -0,0 +1,229 @@ +A method taking execute renders the template when it calls back, with +the data it passes or none. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/flag_custom_names.txtar b/internal/generate/testdata/generate/flag_custom_names.txtar new file mode 100644 index 00000000..ac49c725 --- /dev/null +++ b/internal/generate/testdata/generate/flag_custom_names.txtar @@ -0,0 +1,191 @@ +The output flags rename the generated file and identifiers. + +Command line: 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 +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "Routes", + "ReceiverType": "Server", + "ReceiverPackage": "", + "ReceiverInterface": "Handlers", + "TemplateDataType": "Data", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "Paths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/flag_htmx.txtar b/internal/generate/testdata/generate/flag_htmx.txtar new file mode 100644 index 00000000..eb63cc8f --- /dev/null +++ b/internal/generate/testdata/generate/flag_htmx.txtar @@ -0,0 +1,246 @@ +--output-htmx adds the HX header helpers to TemplateData. + +Command line: muxt generate --output-htmx +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": true, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/flag_logger_path_prefix_middleware.txtar b/internal/generate/testdata/generate/flag_logger_path_prefix_middleware.txtar new file mode 100644 index 00000000..4b2adc3f --- /dev/null +++ b/internal/generate/testdata/generate/flag_logger_path_prefix_middleware.txtar @@ -0,0 +1,219 @@ +The routes function takes a logger, a path prefix, and a middleware. + +Command line: 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 +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": true, + "Logger": true, + "Middleware": true, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/flag_multiple_files.txtar b/internal/generate/testdata/generate/flag_multiple_files.txtar new file mode 100644 index 00000000..f6228cf9 --- /dev/null +++ b/internal/generate/testdata/generate/flag_multiple_files.txtar @@ -0,0 +1,310 @@ +--output-multiple-files writes the routes for each template file into +its own file, beside the main routes file. + +Command line: muxt generate --use-receiver-type=T --output-multiple-files --output-routes-func-with-middleware-param +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": true, + "Verbose": false, + "OutputMultipleFiles": true, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/flag_unexported_identifiers.txtar b/internal/generate/testdata/generate/flag_unexported_identifiers.txtar new file mode 100644 index 00000000..9706973b --- /dev/null +++ b/internal/generate/testdata/generate/flag_unexported_identifiers.txtar @@ -0,0 +1,369 @@ +--output-exported-default-identifiers=false names the generated +identifiers unexported, as the command line spells them. + +Command line: muxt generate --use-receiver-type=Server --output-exported-default-identifiers=false +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "templateRoutes", + "ReceiverType": "Server", + "ReceiverPackage": "", + "ReceiverInterface": "routesReceiver", + "TemplateDataType": "templateData", + "SSETemplateDataType": "sseTemplateData", + "TemplateRoutePathsTypeName": "templateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": false, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- index.gohtml -- +{{define "GET /events sse(Stream(ctx, execute))"}}{{.Result}}{{end}} +{{define "GET /{$}"}}home{{end}} +-- server.go -- +package server + +import "context" + +type Server struct{} + +func (Server) Stream(ctx context.Context, execute func(data string) error) {} +-- want/template_routes.go -- +package server + +import ( + "bytes" + "cmp" + "context" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "path" + "strconv" + "strings" + "sync" +) + +type routesReceiver interface { + Stream(ctx context.Context, execute func(data string) 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() + 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.Stream(ctx, 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(Stream(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())) + return err + } + td.data = buf + mut.Lock() + defer mut.Unlock() + if _, err := td.WriteTo(response); err != nil { + return err + } + flusher.Flush() + return 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 "" +} + +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) Stream() string { + return path.Join(cmp.Or(routePaths.pathsPrefix, "/"), "events") +} + +func (routePaths templateRoutePaths) ReadExact() string { + return "/" +} diff --git a/internal/generate/testdata/generate/flag_without_muxt_version.txtar b/internal/generate/testdata/generate/flag_without_muxt_version.txtar new file mode 100644 index 00000000..70ac0bce --- /dev/null +++ b/internal/generate/testdata/generate/flag_without_muxt_version.txtar @@ -0,0 +1,166 @@ +With --output-muxt-version=false, TemplateData has no MuxtVersion method, +so a template calling it does not compile. + +Command line: muxt generate --output-muxt-version=false +-- config.json -- +{ + "MuxtVersion": "", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": false, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/form_struct.txtar b/internal/generate/testdata/generate/form_struct.txtar new file mode 100644 index 00000000..fae889f9 --- /dev/null +++ b/internal/generate/testdata/generate/form_struct.txtar @@ -0,0 +1,260 @@ +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. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/form_values.txtar b/internal/generate/testdata/generate/form_values.txtar new file mode 100644 index 00000000..4a9f040c --- /dev/null +++ b/internal/generate/testdata/generate/form_values.txtar @@ -0,0 +1,223 @@ +A url.Values parameter receives the parsed form as it is. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/inferred_methods.txtar b/internal/generate/testdata/generate/inferred_methods.txtar new file mode 100644 index 00000000..3a376885 --- /dev/null +++ b/internal/generate/testdata/generate/inferred_methods.txtar @@ -0,0 +1,252 @@ +Without --use-receiver-type every called method is inferred from its +arguments, and returns any. + +Command line: muxt generate +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/last_event_id.txtar b/internal/generate/testdata/generate/last_event_id.txtar new file mode 100644 index 00000000..ffe4c5f2 --- /dev/null +++ b/internal/generate/testdata/generate/last_event_id.txtar @@ -0,0 +1,186 @@ +lastEventID reads the Last-Event-Id header, parsed into the parameter +type. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/marshal_json.txtar b/internal/generate/testdata/generate/marshal_json.txtar new file mode 100644 index 00000000..b9a7ea9a --- /dev/null +++ b/internal/generate/testdata/generate/marshal_json.txtar @@ -0,0 +1,203 @@ +marshalJSON writes the result as JSON, rendering the template for its +side effects and on errors. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/multipart.txtar b/internal/generate/testdata/generate/multipart.txtar new file mode 100644 index 00000000..2dc5b48f --- /dev/null +++ b/internal/generate/testdata/generate/multipart.txtar @@ -0,0 +1,241 @@ +A multipart struct binds text fields and file headers; a +*multipart.Form parameter receives the parsed form. + +Command line: muxt generate --use-receiver-type=T --output-multipart-max-memory=1MiB +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 1048576, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/nested_calls.txtar b/internal/generate/testdata/generate/nested_calls.txtar new file mode 100644 index 00000000..7ba69f84 --- /dev/null +++ b/internal/generate/testdata/generate/nested_calls.txtar @@ -0,0 +1,240 @@ +A call's argument may be the result of another call, a receiver method +or a package function, which runs first. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/path_parameter_types.txtar b/internal/generate/testdata/generate/path_parameter_types.txtar new file mode 100644 index 00000000..755bada0 --- /dev/null +++ b/internal/generate/testdata/generate/path_parameter_types.txtar @@ -0,0 +1,443 @@ +Each path parameter parses into the type of the parameter it is passed +to, and the route path helper takes that type. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/receiver_method_sets.txtar b/internal/generate/testdata/generate/receiver_method_sets.txtar new file mode 100644 index 00000000..844e5aa8 --- /dev/null +++ b/internal/generate/testdata/generate/receiver_method_sets.txtar @@ -0,0 +1,250 @@ +Methods are found through embedded fields and pointer receivers; +functions in the package are called directly and never join the +receiver interface. + +Command line: muxt generate --use-receiver-type=Server +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "Server", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/redirect.txtar b/internal/generate/testdata/generate/redirect.txtar new file mode 100644 index 00000000..d1bccde4 --- /dev/null +++ b/internal/generate/testdata/generate/redirect.txtar @@ -0,0 +1,250 @@ +Only a template that may call a redirect method gets the redirect block. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/request_body.txtar b/internal/generate/testdata/generate/request_body.txtar new file mode 100644 index 00000000..3f1d02e5 --- /dev/null +++ b/internal/generate/testdata/generate/request_body.txtar @@ -0,0 +1,230 @@ +The body is passed as an io.Reader, or decoded as JSON into the +parameter's type. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/response_argument.txtar b/internal/generate/testdata/generate/response_argument.txtar new file mode 100644 index 00000000..12fc8c22 --- /dev/null +++ b/internal/generate/testdata/generate/response_argument.txtar @@ -0,0 +1,172 @@ +A method taking the response writes it itself: muxt writes no status. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/result_shapes.txtar b/internal/generate/testdata/generate/result_shapes.txtar new file mode 100644 index 00000000..665ba3a3 --- /dev/null +++ b/internal/generate/testdata/generate/result_shapes.txtar @@ -0,0 +1,286 @@ +A method returns a value, a value and an error, or a value and a bool. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/route_without_call.txtar b/internal/generate/testdata/generate/route_without_call.txtar new file mode 100644 index 00000000..b24ba8b7 --- /dev/null +++ b/internal/generate/testdata/generate/route_without_call.txtar @@ -0,0 +1,195 @@ +A route with no call renders its template with no result. + +Command line: muxt generate +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/sse.txtar b/internal/generate/testdata/generate/sse.txtar new file mode 100644 index 00000000..270101c0 --- /dev/null +++ b/internal/generate/testdata/generate/sse.txtar @@ -0,0 +1,372 @@ +An sse route streams an event per callback, from the route template or +a same-named one, and reads lastEventID. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/sse_datastar.txtar b/internal/generate/testdata/generate/sse_datastar.txtar new file mode 100644 index 00000000..991f5a1c --- /dev/null +++ b/internal/generate/testdata/generate/sse_datastar.txtar @@ -0,0 +1,408 @@ +With --output-datastar an sse route frames datastar patch events, and +signals decodes the request body. + +Command line: muxt generate --use-receiver-type=T --output-datastar +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": true, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/sse_messages_without_datastar.txtar b/internal/generate/testdata/generate/sse_messages_without_datastar.txtar new file mode 100644 index 00000000..cc0e6a1b --- /dev/null +++ b/internal/generate/testdata/generate/sse_messages_without_datastar.txtar @@ -0,0 +1,344 @@ +An sse method returning nothing, with a data-less callback. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/status_codes.txtar b/internal/generate/testdata/generate/status_codes.txtar new file mode 100644 index 00000000..1b62088e --- /dev/null +++ b/internal/generate/testdata/generate/status_codes.txtar @@ -0,0 +1,277 @@ +The status a route writes: one named in the template name, one the +result reports with a StatusCode method or field, or 200. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/generate/synthesized_method_note.txtar b/internal/generate/testdata/generate/synthesized_method_note.txtar new file mode 100644 index 00000000..87573fe3 --- /dev/null +++ b/internal/generate/testdata/generate/synthesized_method_note.txtar @@ -0,0 +1,229 @@ +With --use-receiver-type, a method the receiver does not define is +inferred, and generation says so. + +Command line: muxt generate --use-receiver-type=T +-- config.json -- +{ + "MuxtVersion": "v1.2.3", + "PackageName": "main", + "PackagePath": "", + "RoutesFunction": "TemplateRoutes", + "ReceiverType": "T", + "ReceiverPackage": "", + "ReceiverInterface": "RoutesReceiver", + "TemplateDataType": "TemplateData", + "SSETemplateDataType": "SSETemplateData", + "TemplateRoutePathsTypeName": "TemplateRoutePaths", + "TemplatesVariables": [ + "templates" + ], + "OutputFileName": "template_routes.go", + "PathPrefix": false, + "Logger": false, + "Middleware": false, + "Verbose": false, + "OutputMultipleFiles": false, + "OutputHTMX": false, + "OutputDatastar": false, + "OutputExportedDefaultIdentifiers": true, + "OutputMuxtVersion": true, + "MultipartMaxMemory": 0, + "SilenceHTTPResponseWarning": false +} +-- 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/load/hydrate_test.go b/internal/load/hydrate_test.go new file mode 100644 index 00000000..4040d758 --- /dev/null +++ b/internal/load/hydrate_test.go @@ -0,0 +1,80 @@ +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") + }) + + 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..64785057 --- /dev/null +++ b/internal/load/loadtest/loadtest.go @@ -0,0 +1,228 @@ +// Package loadtest builds the result of a package load, type checked against +// the official standard library, for tests of what internal/load and its +// callers do with a loaded package. +// +// packages.Load runs go list over the whole package graph, which takes +// seconds. What muxt reads from its result -- syntax, type information, the +// embedded files -- is plain data. Package type checks the files it is given +// in memory and imports the standard library from the go command's export +// data, which the build cache keeps: listing it costs a test binary a couple +// of hundred milliseconds once, and reading a package is shared after that. The +// standard library is whichever the go command in use provides, so the +// tests follow the toolchain rather than a copy of its API. +// +// Package writes the files to disk too, because evaluating ParseFS and +// reading template sources both read them from there. +package loadtest + +import ( + "fmt" + "go/ast" + "go/parser" + "go/token" + "go/types" + "os" + "os/exec" + "path" + "path/filepath" + "slices" + "strings" + "sync" + "testing" + + "golang.org/x/tools/go/gcexportdata" + "golang.org/x/tools/go/packages" +) + +// FileSet positions every package Package loads, and the standard library +// packages imported for them. +var FileSet = token.NewFileSet() + +// The standard library is read from the export data the go command writes +// for it, with gcexportdata -- the reader go/packages uses -- into one map +// shared by every package this test binary loads, so a type a test names is +// the same object wherever it is reached. +// +// Export data lists only the packages a package's API mentions, while a +// load reports every package it imports; check finds fmt, whose functions a +// template calls, through html/template's imports. So each package's +// imports are set to the ones go list reports, as go/packages sets them. +var ( + importLock sync.Mutex + stdlib = make(map[string]*types.Package) + imported = make(map[string]bool) + + stdList = sync.OnceValues(func() (map[string]stdPackage, error) { + // Tab separated: an export path holds whatever the build cache's + // directory is called, spaces included. + out, err := exec.Command("go", "list", "-export", "-f", "{{.ImportPath}}\t{{.Export}}\t{{join .Imports \",\"}}", "std").Output() + if err != nil { + if exitErr, ok := err.(*exec.ExitError); ok { + return nil, fmt.Errorf("go list -export std: %w: %s", err, exitErr.Stderr) + } + return nil, fmt.Errorf("go list -export std: %w", err) + } + list := make(map[string]stdPackage) + for line := range strings.Lines(string(out)) { + fields := strings.Split(strings.TrimRight(line, "\n"), "\t") + // A package the go command builds nothing for -- the + // standard library's own test packages -- has no export + // data to read, so it is not one to import. + if len(fields) < 2 || fields[1] == "" { + continue + } + pkg := stdPackage{export: fields[1]} + if len(fields) > 2 && fields[2] != "" { + pkg.imports = strings.Split(fields[2], ",") + } + list[fields[0]] = pkg + } + return list, nil + }) +) + +type stdPackage struct { + export string + imports []string +} + +func importStd(path string) (*types.Package, error) { + importLock.Lock() + defer importLock.Unlock() + return importStdLocked(path) +} + +func importStdLocked(path string) (*types.Package, error) { + if path == "unsafe" { + return types.Unsafe, nil + } + if pkg, ok := stdlib[path]; ok && imported[path] { + return pkg, nil + } + list, err := stdList() + if err != nil { + return nil, err + } + entry, ok := list[path] + if !ok { + return nil, fmt.Errorf("loadtest imports only the standard library, and %q is not in it", path) + } + f, err := os.Open(entry.export) + if err != nil { + return nil, err + } + defer func() { _ = f.Close() }() + r, err := gcexportdata.NewReader(f) + if err != nil { + return nil, fmt.Errorf("reading export data for %s: %w", path, err) + } + pkg, err := gcexportdata.Read(r, FileSet, stdlib, path) + if err != nil { + return nil, fmt.Errorf("reading export data for %s: %w", path, err) + } + imported[path] = true + imports := make([]*types.Package, 0, len(entry.imports)) + for _, importPath := range entry.imports { + if _, listed := list[importPath]; !listed && importPath != "unsafe" { + continue + } + dep, err := importStdLocked(importPath) + if err != nil { + return nil, err + } + imports = append(imports, dep) + } + pkg.SetImports(imports) + return pkg, nil +} + +type importerFunc func(path string) (*types.Package, error) + +func (fn importerFunc) Import(path string) (*types.Package, error) { return fn(path) } + +// 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 standard library. 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() + 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"): + goPaths = append(goPaths, file) + default: + embedded = append(embedded, file) + } + } + slices.Sort(goPaths) + slices.Sort(embedded) + + // 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. They are imported before the package is + // checked: a package the checker reaches only through another's imports + // is read shallowly, and it is these complete ones -- fmt, whose + // functions a template calls -- that the load's other packages share. + roots := []string{"encoding", "fmt", "net/http"} + std := make([]*types.Package, 0, len(roots)) + for _, stdPath := range roots { + pkg, err := importStd(stdPath) + if err != nil { + t.Fatal(err) + } + std = append(std, pkg) + } + + syntax := make([]*ast.File, 0, len(goPaths)) + for _, file := range goPaths { + parsed, err := parser.ParseFile(FileSet, file, files[filepath.Base(file)], parser.ParseComments|parser.SkipObjectResolution) + if err != nil { + t.Fatal(err) + } + syntax = append(syntax, parsed) + } + 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), + } + config := types.Config{Importer: importerFunc(importStd)} + pkg, err := config.Check(pkgPath, FileSet, syntax, info) + if err != nil { + t.Fatal(err) + } + + pl := []*packages.Package{{ + ID: pkgPath, + Name: pkg.Name(), + PkgPath: pkgPath, + Dir: dir, + GoFiles: goPaths, + EmbedFiles: embedded, + Fset: FileSet, + Syntax: syntax, + Types: pkg, + TypesInfo: info, + }} + for _, pkg := range std { + pl = append(pl, &packages.Package{ID: pkg.Path(), Name: pkg.Name(), PkgPath: pkg.Path(), Fset: FileSet, Types: pkg}) + } + 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/load/package.go b/internal/load/package.go index c3ad1684..8d9ecd05 100644 --- a/internal/load/package.go +++ b/internal/load/package.go @@ -122,31 +122,6 @@ func PackageInDirectory(list []*packages.Package, dir 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 Templates(wd, templatesVariable string, pl []*packages.Package) (*LoadedTemplates, error) { - pkg, ok := PackageInDirectory(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 -} - // 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/load/source.go b/internal/load/source.go index c669c90a..7b97b13d 100644 --- a/internal/load/source.go +++ b/internal/load/source.go @@ -22,9 +22,8 @@ func Package(dir string, pl []*packages.Package, variables []string) (source.Pac return source.Package{}, NoPackageError(dir, pl) } result := source.Package{ - Fset: pkg.Fset, - Types: pkg.Types, - Imports: imports(pl), + Fset: pkg.Fset, + Types: pkg.Types, } for _, name := range variables { variable, err := Variable(pkg, name) @@ -83,32 +82,3 @@ func Receiver(dir string, pl []*packages.Package, packagePath, ident string) (*t } 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/load/stdlib.go b/internal/load/stdlib.go new file mode 100644 index 00000000..884287a2 --- /dev/null +++ b/internal/load/stdlib.go @@ -0,0 +1,112 @@ +package load + +import ( + "fmt" + "go/types" + + "golang.org/x/tools/go/packages" + + "github.com/typelate/muxt/internal/muxt" +) + +// StandardLibrary answers route resolution's questions about the standard +// library from the packages a load reached: those in pl and every package +// they import. It is the one muxt.Checker backed by the official standard +// library, whichever version the go command loaded. +func StandardLibrary(pl []*packages.Package) muxt.Checker { + return standardLibrary(indexImports(pl)) +} + +// standardLibrary is a muxt.Checker over packages indexed by import path. +type standardLibrary map[string]*types.Package + +// scopeTypes names, for each reserved argument identifier, the standard +// library type it binds to and whether it binds a pointer to that type. +var scopeTypes = map[string]struct { + path, name string + pointer bool +}{ + muxt.TemplateNameScopeIdentifierHTTPRequest: {"net/http", "Request", true}, + muxt.TemplateNameScopeIdentifierHTTPResponse: {"net/http", "ResponseWriter", false}, + muxt.TemplateNameScopeIdentifierContext: {"context", "Context", false}, + muxt.TemplateNameScopeIdentifierForm: {"net/url", "Values", false}, + muxt.TemplateNameScopeIdentifierMultipart: {"mime/multipart", "Form", true}, + muxt.TemplateNameScopeIdentifierRequestBody: {"io", "Reader", false}, +} + +func (std standardLibrary) ScopeType(identifier string) (types.Type, error) { + scope, ok := scopeTypes[identifier] + if !ok { + return nil, fmt.Errorf("%s is not a reserved argument identifier", identifier) + } + return std.lookup(scope.path, scope.name, scope.pointer) +} + +func (std standardLibrary) FileHeader() (types.Type, error) { + return std.lookup("mime/multipart", "FileHeader", true) +} + +func (std standardLibrary) RawJSON() (types.Type, error) { + return std.lookup("encoding/json", "RawMessage", false) +} + +func (std standardLibrary) TextUnmarshaler(tp types.Type) bool { + return std.implements(types.NewPointer(tp), "TextUnmarshaler") +} + +func (std standardLibrary) TextMarshaler(tp types.Type) bool { + return std.implements(tp, "TextMarshaler") +} + +func (std standardLibrary) implements(tp types.Type, encodingInterface string) bool { + iface, err := std.lookup("encoding", encodingInterface, false) + if err != nil { + return false + } + return types.Implements(tp, iface.Underlying().(*types.Interface)) +} + +func (std standardLibrary) lookup(path, name string, pointer bool) (types.Type, error) { + pkg, ok := std[path] + if !ok { + return nil, fmt.Errorf("could not find package %q for %s", path, name) + } + obj := pkg.Scope().Lookup(name) + if obj == nil { + return nil, fmt.Errorf("package %q declares no %s", path, name) + } + tp := obj.Type() + if pointer { + tp = types.NewPointer(tp) + } + return tp, nil +} + +// indexImports 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 indexImports(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/load/stdlib_test.go b/internal/load/stdlib_test.go new file mode 100644 index 00000000..3d29eb85 --- /dev/null +++ b/internal/load/stdlib_test.go @@ -0,0 +1,86 @@ +package load_test + +import ( + "go/types" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/load" + "github.com/typelate/muxt/internal/load/loadtest" + "github.com/typelate/muxt/internal/muxt" +) + +// TestStandardLibrary states what the one checker backed by the official +// standard library answers, which the mock checkers in other packages' tests +// stand in for. +func TestStandardLibrary(t *testing.T) { + dir := t.TempDir() + pl := loadtest.Package(t, dir, "example.com/server", map[string]string{ + "server.go": `package server + +import ( + "html/template" + "time" +) + +type Plain struct{} + +var ( + _ time.Time + _ template.HTML +) +`, + }) + std := load.StandardLibrary(pl) + + t.Run("the reserved identifiers", func(t *testing.T) { + for identifier, want := range map[string]string{ + muxt.TemplateNameScopeIdentifierHTTPRequest: "*net/http.Request", + muxt.TemplateNameScopeIdentifierHTTPResponse: "net/http.ResponseWriter", + muxt.TemplateNameScopeIdentifierContext: "context.Context", + muxt.TemplateNameScopeIdentifierForm: "net/url.Values", + muxt.TemplateNameScopeIdentifierMultipart: "*mime/multipart.Form", + muxt.TemplateNameScopeIdentifierRequestBody: "io.Reader", + } { + tp, err := std.ScopeType(identifier) + require.NoError(t, err, identifier) + assert.Equal(t, want, types.TypeString(tp, nil), identifier) + } + _, err := std.ScopeType("lastEventID") + assert.EqualError(t, err, "lastEventID is not a reserved argument identifier") + }) + + t.Run("the types a binding needs", func(t *testing.T) { + fileHeader, err := std.FileHeader() + require.NoError(t, err) + assert.Equal(t, "*mime/multipart.FileHeader", types.TypeString(fileHeader, nil)) + rawJSON, err := std.RawJSON() + require.NoError(t, err) + assert.Equal(t, "encoding/json.RawMessage", types.TypeString(rawJSON, nil)) + }) + + t.Run("text marshaling", func(t *testing.T) { + var timeType types.Type + for _, imported := range pl[0].Types.Imports() { + if imported.Path() == "time" { + timeType = imported.Scope().Lookup("Time").Type() + } + } + require.NotNil(t, timeType) + plain := pl[0].Types.Scope().Lookup("Plain").Type() + assert.True(t, std.TextUnmarshaler(timeType), "a *time.Time parses from text") + assert.True(t, std.TextMarshaler(timeType), "a time.Time formats as text") + assert.False(t, std.TextUnmarshaler(plain)) + assert.False(t, std.TextMarshaler(plain)) + }) + + t.Run("a package the load did not reach", func(t *testing.T) { + // Only the package itself, which imports nothing: encoding/json is + // reached through html/template's imports in a real package. + bare := t.TempDir() + _, err := load.StandardLibrary(loadtest.Package(t, bare, "example.com/bare", map[string]string{"bare.go": "package bare\n"})[:1]).RawJSON() + assert.EqualError(t, err, `could not find package "encoding/json" for RawMessage`) + }) +} diff --git a/internal/load/templates_test.go b/internal/load/templates_test.go index 1c7d345e..b5bae801 100644 --- a/internal/load/templates_test.go +++ b/internal/load/templates_test.go @@ -140,9 +140,8 @@ func TestPackageInADirectoryNamedLikeAGoFile(t *testing.T) { require.NoError(t, err) assert.Equal(t, "scratch", pkg.Types.Path()) - // The mutation run still reads Templates, and it looks the package up - // the same way. - lt, err := load.Templates(dir, "templates", pl) - require.NoError(t, err) - require.NotNil(t, lt.HTML.Lookup("home")) + // Every command reads the package through Package now, the mutation + // run included. + require.Len(t, pkg.Variables, 1) + require.NotNil(t, pkg.Variables[0].Set.Lookup("home")) } diff --git a/internal/mutation/collect.go b/internal/mutation/collect.go index 4a9ad14a..c4a67084 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,11 +20,14 @@ 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 + // literals are where the string literals of each Go file sit, read + // once per file: a run asks for every definition written in it. + literals map[string][]literalSpan + // delims are the delimiters each source was parsed with, read off the // definitions before any source is built. A source scans its actions // as it is constructed, so the delimiters have to be known by then, @@ -47,7 +49,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,23 +76,38 @@ 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 := c.literalAt(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 { +// literalAt returns the bounds of the Go string literal covering offset in +// file, reading the file's literals the first time it is asked about it. +func (c *sourceCollector) literalAt(file, text string, offset int) (start, end int, ok bool) { + spans, read := c.literals[file] + if !read { + spans = stringLiterals(file, text) + c.literals[file] = spans + } + return literalAt(spans, offset) +} + +func newSourceCollector(workingDirectory string, defs []source.Definition) *sourceCollector { c := &sourceCollector{ workingDirectory: workingDirectory, - packages: pl, files: make(map[string]string), + literals: make(map[string][]literalSpan), byKey: make(map[sourceKey]*templateSource), delims: make(map[sourceKey][2]string), } @@ -108,10 +125,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 +148,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 := c.literalAt(file, fileText, definition.Define.Offset) if !ok { return nil, nil } diff --git a/internal/mutation/configuration_test.go b/internal/mutation/configuration_test.go new file mode 100644 index 00000000..eb42fe5f --- /dev/null +++ b/internal/mutation/configuration_test.go @@ -0,0 +1,53 @@ +package mutation_test + +import ( + "encoding/json/v2" + "regexp" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/configjson" + "github.com/typelate/muxt/internal/mutation" +) + +// TestConfigurationJSON states how a run's configuration reads and writes +// as JSON, which is how an archive in testdata holds the one it plans with: +// a pattern is the text the command line wrote, and a flag the command line +// left alone is null. A pattern read as "" would be a pattern matching +// everything, which is not what leaving --run alone means, so the archives +// hold null. +func TestConfigurationJSON(t *testing.T) { + for _, tt := range []struct { + name string + config mutation.Configuration + }{ + { + name: "no patterns", + config: mutation.Configuration{TemplatesVariables: []string{"templates"}, Packages: []string{}, DryRun: true, Seed: 1, SeedSet: true, MaxCases: mutation.DefaultMaxCases, Workers: 1}, + }, + { + name: "a pattern is the text it was written as", + config: mutation.Configuration{TemplatesVariables: []string{"templates"}, TemplatePattern: regexp.MustCompile("^footer$"), Run: regexp.MustCompile("TestPage"), Packages: []string{"./..."}, MaxCases: 2, Workers: 4, Diff: "main", GoTestArgs: []string{"-count=1"}}, + }, + } { + t.Run(tt.name, func(t *testing.T) { + written, err := json.Marshal(tt.config, configjson.Options()) + require.NoError(t, err) + + var read mutation.Configuration + require.NoError(t, json.Unmarshal(written, &read, configjson.Options())) + assert.Equal(t, tt.config, read) + }) + } +} + +// TestConfigurationJSONRejectsABadPattern states that a pattern that does +// not compile is reported where it was read. +func TestConfigurationJSONRejectsABadPattern(t *testing.T) { + var config mutation.Configuration + err := json.Unmarshal([]byte(`{"TemplatePattern":"("}`), &config, configjson.Options()) + require.ErrorContains(t, err, "error parsing regexp") + require.ErrorContains(t, err, "TemplatePattern") +} 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 a668f728..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/load" ) // 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 := load.Templates(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..e50f5ec6 --- /dev/null +++ b/internal/mutation/input.go @@ -0,0 +1,73 @@ +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 this and the +// template files its definitions name, and needs no go command, 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 + 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, + 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 ec7559fc..d249ed1b 100644 --- a/internal/mutation/plan.go +++ b/internal/mutation/plan.go @@ -10,10 +10,9 @@ 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/load" + "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 := load.Templates(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 *load.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 *load.LoadedTemplates, sc scope, functions check.Functions // 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 *load.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 *load.LoadedTemplates, sc scope, mutant Mutant, functions check. // 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 *load.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,18 +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 := load.PackagesWithEnv(workingDirectory, env) - return pl, err - } - return load.PackagesWithTests(workingDirectory, env) -} - func relativePosition(workingDirectory string, position token.Position) string { if !position.IsValid() { return "?" diff --git a/internal/mutation/report.go b/internal/mutation/report.go index f700342f..8ba304ec 100644 --- a/internal/mutation/report.go +++ b/internal/mutation/report.go @@ -44,7 +44,7 @@ type Result struct { Reason string `json:"reason,omitempty"` // Seconds is how long the mutant's test run took. - Seconds float64 `json:"seconds,omitempty"` + Seconds float64 `json:"seconds,omitzero"` // mutantIndex locates the mutant this result came from, so the run // does not have to carry the mutants inside the report it prints. @@ -107,7 +107,7 @@ type BaselineResult struct { // Seconds is how long the unmutated run took, which is what the // estimate for the whole run is built from. - Seconds float64 `json:"seconds,omitempty"` + Seconds float64 `json:"seconds,omitzero"` } // Report is the outcome of a whole mutation run. diff --git a/internal/mutation/runner.go b/internal/mutation/runner.go index 5175949f..405a4e34 100644 --- a/internal/mutation/runner.go +++ b/internal/mutation/runner.go @@ -1,7 +1,7 @@ package mutation import ( - "encoding/json" + "encoding/json/v2" "fmt" "io" "os" diff --git a/internal/mutation/runner_test.go b/internal/mutation/runner_test.go index fe84e741..dc4ccb10 100644 --- a/internal/mutation/runner_test.go +++ b/internal/mutation/runner_test.go @@ -1,7 +1,7 @@ package mutation import ( - "encoding/json" + "encoding/json/v2" "errors" "fmt" "os" diff --git a/internal/mutation/snapshot_test.go b/internal/mutation/snapshot_test.go new file mode 100644 index 00000000..7a4f625e --- /dev/null +++ b/internal/mutation/snapshot_test.go @@ -0,0 +1,191 @@ +package mutation + +import ( + "encoding/json/v2" + "errors" + "flag" + "os" + "path/filepath" + "slices" + "strings" + "testing" + + "golang.org/x/tools/txtar" + + "github.com/typelate/muxt/internal/configjson" + "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/test-template-mutations and compares the report with the +// archive's want/ files. The directory is the command, so one case runs +// with -run TestSnapshots/test-template-mutations/diff. +// +// An archive holds everything a case is: the configuration it plans with, +// its inputs, and what the plan reports. +// +// - config.json is the configuration, as the command line named in the +// archive's header parses into it. TestCommandLineConfigurations in +// internal/cli states that parse; this states what planning does with +// the result. +// - Go files and template files are written to a directory and loaded as +// example.com/server by internal/load/loadtest: type checked against +// the official standard library, without loading the package graph. +// - 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) { + const command = "test-template-mutations" + if stray, _ := filepath.Glob(filepath.Join("testdata", "*.txtar")); len(stray) > 0 { + t.Fatalf("%s is not in a command's directory; the mutation run's archives are in testdata/%s", stray[0], command) + } + archives, err := filepath.Glob(filepath.Join("testdata", command, "*.txtar")) + if err != nil { + t.Fatal(err) + } + if len(archives) == 0 { + t.Fatalf("no archives in testdata/%s", command) + } + t.Run(command, func(t *testing.T) { + for _, archivePath := range archives { + runSnapshot(t, archivePath) + } + }) +} + +// runSnapshot compares one archive's want/ files with the report planning +// it produces, or with -update rewrites them. +func runSnapshot(t *testing.T, archivePath string) { + t.Helper() + t.Run(strings.TrimSuffix(filepath.Base(archivePath), ".txtar"), func(t *testing.T) { + archive, err := txtar.ParseFile(archivePath) + if err != nil { + t.Fatal(err) + } + config := configuration(t, archive) + if !config.DryRun { + t.Fatal("a snapshot plans a dry run; config.json must set DryRun") + } + got := dryRunSnapshot(t, 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]) + } + } + }) +} + +// configuration reads the archive's config.json: the configuration to plan +// with. +func configuration(t *testing.T, archive *txtar.Archive) Configuration { + t.Helper() + for _, file := range archive.Files { + if file.Name != "config.json" { + continue + } + var config Configuration + if err := json.Unmarshal(file.Data, &config, configjson.Options()); err != nil { + t.Fatalf("config.json: %v", err) + } + return config + } + t.Fatal("the archive has no config.json") + return Configuration{} +} + +// 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 file.Name == "config.json": + 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/source.go b/internal/mutation/source.go index a6513210..1c40916e 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,38 +254,44 @@ 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 - } +// literalSpan is where one string literal sits in its file's text. +type literalSpan struct{ start, end int } + +// stringLiterals returns where every string literal in the file sits, in +// the order they were written, which is how a template written in Go +// source is found in its file's text. +// +// The file is parsed here rather than taken from the loaded package: the +// text is already in hand, and it leaves the plan needing nothing from the +// loader. AllErrors keeps a file that does not fully parse yielding the +// literals the parser did read -- without it the parser gives up after ten +// syntax errors and returns nothing, and a run over a broken file would +// report its templates as unreadable rather than the file as unparsed. +func stringLiterals(filename, text string) []literalSpan { + fset := token.NewFileSet() + file, _ := parser.ParseFile(fset, filename, text, parser.SkipObjectResolution|parser.AllErrors) + if file == nil { + return nil + } + tokenFile := fset.File(file.FileStart) + var spans []literalSpan + ast.Inspect(file, func(node ast.Node) bool { + lit, isLit := node.(*ast.BasicLit) + if !isLit || lit.Kind != token.STRING { + return true + } + spans = append(spans, literalSpan{start: tokenFile.Offset(lit.Pos()), end: tokenFile.Offset(lit.End())}) + return false + }) + return spans +} + +// literalAt returns the bounds of the literal covering offset. Literals do +// not nest, so the first match is the one wanted. +func literalAt(spans []literalSpan, offset int) (start, end int, ok bool) { + for _, span := range spans { + if offset >= span.start && offset < span.end { + return span.start, span.end, true } } return 0, 0, false diff --git a/internal/mutation/testdata/test-template-mutations/delimiters.txtar b/internal/mutation/testdata/test-template-mutations/delimiters.txtar new file mode 100644 index 00000000..d3540e91 --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/delimiters.txtar @@ -0,0 +1,43 @@ +A literal parsed with other delimiters is read with them. + +Command line: muxt test-template-mutations --dry-run --seed=1 -v +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": true, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/diff.txtar b/internal/mutation/testdata/test-template-mutations/diff.txtar new file mode 100644 index 00000000..93bca4f4 --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/diff.txtar @@ -0,0 +1,85 @@ +With --diff, a template whose text and dot are unchanged since the +revision is left alone and reported; one that reads differently is +mutated. + +Command line: muxt test-template-mutations --dry-run --seed=1 -v --diff=main +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": true, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "main", + "GoTestArgs": null +} +-- 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/test-template-mutations/err_missing_variable.txtar b/internal/mutation/testdata/test-template-mutations/err_missing_variable.txtar new file mode 100644 index 00000000..64cf14f2 --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/err_missing_variable.txtar @@ -0,0 +1,42 @@ +A templates variable the package does not declare. + +Command line: muxt test-template-mutations --dry-run --seed=1 --use-templates-variable=pages +-- config.json -- +{ + "TemplatesVariables": [ + "pages" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": false, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/err_no_call_sites.txtar b/internal/mutation/testdata/test-template-mutations/err_no_call_sites.txtar new file mode 100644 index 00000000..b404b3ce --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/err_no_call_sites.txtar @@ -0,0 +1,42 @@ +A templates variable whose templates no ExecuteTemplate call renders has +nothing to mutate. + +Command line: muxt test-template-mutations --dry-run --seed=1 +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": false, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/err_no_mutations.txtar b/internal/mutation/testdata/test-template-mutations/err_no_mutations.txtar new file mode 100644 index 00000000..58832064 --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/err_no_mutations.txtar @@ -0,0 +1,43 @@ +Templates that were reached but hold no action to vary are an error, not +a pass. + +Command line: muxt test-template-mutations --dry-run --seed=1 +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": false, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/literal_template.txtar b/internal/mutation/testdata/test-template-mutations/literal_template.txtar new file mode 100644 index 00000000..b0a12a8f --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/literal_template.txtar @@ -0,0 +1,51 @@ +A template written as a Go string literal: its actions and branches are +mutated where they are written in the literal. + +Command line: muxt test-template-mutations --dry-run --seed=1 -v +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": true, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/operand_budget.txtar b/internal/mutation/testdata/test-template-mutations/operand_budget.txtar new file mode 100644 index 00000000..8db628e5 --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/operand_budget.txtar @@ -0,0 +1,53 @@ +An action with more operand combinations than --max-cases allows +contributes none, and says so where its mutants would have been. + +Command line: muxt test-template-mutations --dry-run --seed=1 -v --max-cases=2 +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": true, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 2, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/partials_and_trims.txtar b/internal/mutation/testdata/test-template-mutations/partials_and_trims.txtar new file mode 100644 index 00000000..b2a4bc9f --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/partials_and_trims.txtar @@ -0,0 +1,77 @@ +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. + +Command line: muxt test-template-mutations --dry-run --seed=1 -v +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": true, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/skipped_mutants.txtar b/internal/mutation/testdata/test-template-mutations/skipped_mutants.txtar new file mode 100644 index 00000000..c1543e34 --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/skipped_mutants.txtar @@ -0,0 +1,65 @@ +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. + +Command line: muxt test-template-mutations --dry-run --seed=1 -v +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": null, + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": true, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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/test-template-mutations/template_pattern.txtar b/internal/mutation/testdata/test-template-mutations/template_pattern.txtar new file mode 100644 index 00000000..964ba793 --- /dev/null +++ b/internal/mutation/testdata/test-template-mutations/template_pattern.txtar @@ -0,0 +1,53 @@ +--template-pattern limits mutation to the templates it matches; the +traversal still passes through the others to reach them. + +Command line: muxt test-template-mutations --dry-run --seed=1 -v --template-pattern=^footer$ +-- config.json -- +{ + "TemplatesVariables": [ + "templates" + ], + "TemplatePattern": "^footer$", + "Run": null, + "Packages": [], + "DryRun": true, + "Verbose": true, + "IncludeTests": false, + "Seed": 1, + "SeedSet": true, + "MaxCases": 8, + "Workers": 1, + "Diff": "", + "GoTestArgs": null +} +-- 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"}}
    {{.Title}}
    {{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 6c1e6aae..2a1d2e8b 100644 --- a/internal/mutation/traverse.go +++ b/internal/mutation/traverse.go @@ -6,7 +6,6 @@ import ( "text/template/parse" "github.com/typelate/check" - "github.com/typelate/muxt/internal/load" ) // callSite is one templates.ExecuteTemplate call, which is where a @@ -79,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 *load.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 } @@ -95,7 +94,7 @@ func traverse(lt *load.LoadedTemplates, index map[string]treeLocation) ([]scope, // traversal is the state of one walk: what it reads from, and what it has // found so far. type traversal struct { - lt *load.LoadedTemplates + lt *checked index map[string]treeLocation // visited records the call site each template and dot was first @@ -147,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 *load.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 @@ -160,18 +159,18 @@ func templateCalls(lt *load.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 *load.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 *load.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 eedcc4f0..62a0a2fd 100644 --- a/internal/muxt/call.go +++ b/internal/muxt/call.go @@ -10,6 +10,7 @@ import ( "slices" "strings" + "github.com/typelate/muxt/internal/asteval" "github.com/typelate/muxt/internal/astgen" "github.com/typelate/muxt/internal/source" ) @@ -17,16 +18,17 @@ import ( type Argument struct { Identifier string Type ArgumentType - ParamType types.Type + paramType types.Type template *template.Template // sig, args, and isMethod describe a nested call argument // (Type == ArgumentTypeCall): the nested call's signature, its own // hydrated arguments, and whether it resolves to a receiver method. - sig *types.Signature - args []Argument - isMethod bool + sig *types.Signature + args []Argument + isMethod bool + resultShape ResultShape // callbackResult and callbackHasArg describe a validated render-callback // argument (Type == ArgumentTypeExecute): the template data type T the @@ -39,11 +41,37 @@ type Argument struct { // argument binds to the request (nil for the raw url.Values / // *multipart.Form mode). formFields []FieldBinding + + // scopeType, direct and method describe a request value argument: the + // type the identifier binds to, whether that value is assignable to the + // parameter as it is, and otherwise how it parses from its string form. + scopeType types.Type + direct bool + method UnmarshalMethod } +// ScopeType returns the type a request value argument binds to before any +// parsing: *http.Request for request, string for a path value, and so on. +func (a Argument) ScopeType() source.Type { return source.NewType(a.scopeType) } + +// Direct reports whether a request value argument is assignable to its +// parameter as it is. A path value or lastEventID that is not assignable +// parses with UnmarshalMethod; a form or multipart argument that is not +// assignable binds a struct through FormFields. +func (a Argument) Direct() bool { return a.direct } + +// UnmarshalMethod returns how a path value or lastEventID argument that is +// not Direct parses from its string form. +func (a Argument) UnmarshalMethod() UnmarshalMethod { return a.method } + // Signature returns the resolved signature of a nested call argument // (Type == ArgumentTypeCall), or nil for a leaf argument. -func (a Argument) Signature() *types.Signature { return a.sig } +func (a Argument) Signature() source.Type { + if a.sig == nil { + return source.Type{} + } + return source.NewType(a.sig) +} // IsMethod reports whether a nested call argument resolves to a receiver method // (as opposed to a package-scope function). @@ -52,6 +80,12 @@ func (a Argument) IsMethod() bool { return a.isMethod } // Arguments returns the hydrated arguments of a nested call argument. func (a Argument) Arguments() []Argument { return a.args } +// ParamType is the type of the parameter the argument is passed to. +func (a Argument) ParamType() source.Type { return source.NewType(a.paramType) } + +// ResultShape classifies a nested call argument's results. +func (a Argument) ResultShape() ResultShape { return a.resultShape } + // Template returns the template a render-callback argument (ArgumentTypeExecute) // renders: the route template for the base execute callback, or the same-named // template for an sse-prefixed callback (nil if that template does not exist). @@ -60,16 +94,16 @@ func (a Argument) Template() *template.Template { return a.template } // CallbackSignature returns a render-callback argument's function signature // (from its parameter type), or nil if the parameter type is not a function. func (a Argument) CallbackSignature() *types.Signature { - if a.ParamType == nil { + if a.paramType == nil { return nil } - sig, _ := a.ParamType.Underlying().(*types.Signature) + sig, _ := a.paramType.Underlying().(*types.Signature) return sig } // CallbackResultType returns the template data type T a validated // render-callback argument receives (struct{} for a func() error callback). -func (a Argument) CallbackResultType() types.Type { return a.callbackResult } +func (a Argument) CallbackResultType() source.Type { return source.NewType(a.callbackResult) } // CallbackHasArg reports whether a validated render-callback argument's // callback takes the template data argument (func(T) error vs func() error). @@ -118,18 +152,20 @@ const ( ResultShapeError ) -func ResolveCall(def *Definition, pkg source.Package, receiver *types.Named) error { +// ResolveCall resolves def's call against the receiver type and the package's +// functions, asking checker what it needs to know about the standard library. +func ResolveCall(def *Definition, pkg source.Package, receiver *types.Named, checker Checker) error { if def.call == nil || def.fun == nil { return nil } - sig, isMethod, args, err := resolveCall(def, def.call, pkg, receiver) + sig, isMethod, args, err := resolveCall(def, def.call, pkg, receiver, checker) 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)) + recordPathValueTypes(def, checker, 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 @@ -137,9 +173,58 @@ func ResolveCall(def *Definition, pkg source.Package, receiver *types.Named) err return def.finishNameError(errAtNode(def.fun, err), def.handlerSpan()) } def.resultShape = shape - return def.finishNameError(resolveCallbackShapes(def), def.handlerSpan()) + if err := resolveCallbackShapes(def); err != nil { + return def.finishNameError(err, def.handlerSpan()) + } + def.resultStatusCode = statusCodeSource(def.resultType(), pkg.Types) + return nil +} + +// ResultStatusCode is where a route's result offers a status code: a +// StatusCode() int method, a StatusCode field, or nowhere. +type ResultStatusCode int + +const ( + ResultStatusCodeNone ResultStatusCode = iota + ResultStatusCodeMethod + ResultStatusCodeField +) + +// resultType is the type the template data's Result field has: the +// execute callback's parameter when the call takes one, otherwise the +// call's first result. +func (def *Definition) resultType() types.Type { + if i, ok := def.ExecuteArgumentIndex(); ok { + return def.Arguments[i].callbackResult + } + switch def.resultShape { + case ResultShapeData, ResultShapeDataError, ResultShapeDataOK: + return def.sig.Results().At(0).Type() + } + return nil +} + +func statusCodeSource(tp types.Type, pkg *types.Package) ResultStatusCode { + if tp == nil { + return ResultStatusCodeNone + } + if types.Implements(tp, statusCoder) { + return ResultStatusCodeMethod + } + if obj, _, _ := types.LookupFieldOrMethod(tp, true, pkg, "StatusCode"); obj != nil { + return ResultStatusCodeField + } + return ResultStatusCodeNone } +var statusCoder = types.NewInterfaceType([]*types.Func{ + types.NewFunc(token.NoPos, nil, "StatusCode", types.NewSignatureType(nil, nil, nil, + types.NewTuple(), + types.NewTuple(types.NewVar(token.NoPos, nil, "", types.Typ[types.Int])), + false, + )), +}, nil).Complete() + // recordPathValueTypes records the type each path parameter parses into. // // A parameter is parsed once per request, where the call first passes it @@ -148,18 +233,25 @@ func ResolveCall(def *Definition, pkg source.Package, receiver *types.Named) err // 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) { +func recordPathValueTypes(def *Definition, checker Checker, args []Argument, seen map[string]bool) { for _, arg := range args { switch arg.Type { case ArgumentTypeCall: - recordPathValueTypes(into, arg.args, seen) + recordPathValueTypes(def, checker, 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 + if isStringAssignable(arg.paramType) { + continue + } + for i := range def.Segments { + segment := &def.Segments[i] + if segment.IsWildcard() && segment.value == arg.Identifier { + segment.tp = arg.paramType + segment.textMarshaler = checker.TextMarshaler(arg.paramType) + } } } } @@ -271,28 +363,28 @@ func classifyResultShape(def *Definition, qual types.Qualifier) (ResultShape, er } } -// checkNestedCallResultShape validates a nested call's results: one value, -// optionally followed by an error or bool. -func checkNestedCallResultShape(name string, sig *types.Signature, qual types.Qualifier) error { +// classifyNestedCallResultShape validates a nested call's results: one +// value, optionally followed by an error or bool. +func classifyNestedCallResultShape(name string, sig *types.Signature, qual types.Qualifier) (ResultShape, error) { results := sig.Results() errIface := types.Universe.Lookup("error").Type().Underlying().(*types.Interface) sigStr := name + strings.TrimPrefix(types.TypeString(sig, qual), "func") switch results.Len() { case 1: - return nil + return ResultShapeData, nil case 2: last := results.At(1).Type() if types.Implements(last, errIface) { - return nil + return ResultShapeDataError, nil } if basic, ok := last.(*types.Basic); ok && basic.Kind() == types.Bool { - return nil + return ResultShapeDataOK, nil } - return fmt.Errorf("the second result of %s must be an error or a bool, got %s", sigStr, types.TypeString(last, qual)) + return ResultShapeInvalid, fmt.Errorf("the second result of %s must be an error or a bool, got %s", sigStr, types.TypeString(last, qual)) case 0: - return fmt.Errorf("method %s has no results; it should have one or two", sigStr) + return ResultShapeInvalid, fmt.Errorf("method %s has no results; it should have one or two", sigStr) default: - return fmt.Errorf("method %s has %d results; it should have one or two", sigStr, results.Len()) + return ResultShapeInvalid, fmt.Errorf("method %s has %d results; it should have one or two", sigStr, results.Len()) } } @@ -318,7 +410,7 @@ func definedHere(pkg source.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, pkg source.Package, receiver *types.Named) (*types.Signature, bool, []Argument, error) { +func resolveCall(def *Definition, call *ast.CallExpr, pkg source.Package, receiver *types.Named, checker Checker) (*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)) @@ -330,7 +422,7 @@ func resolveCall(def *Definition, call *ast.CallExpr, pkg source.Package, receiv object = m isMethod = false } else { - ms, err := synthesizeCallSignature(def, call, pkg, receiver) + ms, err := synthesizeCallSignature(def, call, pkg, receiver, checker) if err != nil { return nil, false, nil, err } @@ -379,7 +471,7 @@ func resolveCall(def *Definition, call *ast.CallExpr, pkg source.Package, receiv args = append(args, Argument{Identifier: argument.Name}) continue } - arg, err := newArgumentFromIdentifier(def, pkg, argument, paramType, qual) + arg, err := newArgumentFromIdentifier(def, checker, argument, paramType, qual) if err != nil { return nil, false, nil, errAtNode(argument, err) } @@ -400,24 +492,26 @@ func resolveCall(def *Definition, call *ast.CallExpr, pkg source.Package, receiv args = append(args, Argument{ Identifier: TemplateNameScopeIdentifierRequestBody, Type: ArgumentTypeRequestBodyJSON, - ParamType: paramType, + paramType: paramType, }) continue } - nestedSig, nestedIsMethod, nestedArgs, err := resolveCall(def, argument, pkg, receiver) + nestedSig, nestedIsMethod, nestedArgs, err := resolveCall(def, argument, pkg, receiver, checker) if err != nil { return nil, false, nil, err } - if err := checkNestedCallResultShape(name, nestedSig, qual); err != nil { + nestedShape, err := classifyNestedCallResultShape(name, nestedSig, qual) + if err != nil { return nil, false, nil, errAtNode(argument.Fun, err) } args = append(args, Argument{ - Identifier: name, - Type: ArgumentTypeCall, - ParamType: paramType, - sig: nestedSig, - isMethod: nestedIsMethod, - args: nestedArgs, + Identifier: name, + Type: ArgumentTypeCall, + paramType: paramType, + sig: nestedSig, + isMethod: nestedIsMethod, + args: nestedArgs, + resultShape: nestedShape, }) } } @@ -428,7 +522,7 @@ func resolveCall(def *Definition, call *ast.CallExpr, pkg source.Package, receiv // 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, pkg source.Package, receiver *types.Named) (*types.Signature, error) { +func synthesizeCallSignature(def *Definition, call *ast.CallExpr, pkg source.Package, receiver *types.Named, checker Checker) (*types.Signature, error) { var params []*types.Var hasSSE := false // Each argument becomes a parameter named after it, so a repeated @@ -461,7 +555,7 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, pkg source.Pac } continue } - tp, ok := DefaultScopeType(pkg, def, arg.Name) + tp, ok := defaultScopeType(checker, def, arg.Name) if !ok { return nil, errAt(arg, "could not determine a type for %s", arg.Name) } @@ -472,7 +566,7 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, pkg source.Pac 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(pkg, "encoding/json", "RawMessage", false) + tp, err := checker.RawJSON() if err != nil { return nil, err } @@ -481,7 +575,7 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, pkg source.Pac } continue } - if _, _, _, err := resolveCall(def, arg, pkg, receiver); err != nil { + if _, _, _, err := resolveCall(def, arg, pkg, receiver, checker); err != nil { return nil, err } } @@ -493,35 +587,19 @@ func synthesizeCallSignature(def *Definition, call *ast.CallExpr, pkg source.Pac return types.NewSignatureType(types.NewVar(0, nil, "", receiver.Obj().Type()), nil, nil, types.NewTuple(params...), results, false), nil } -func DefaultScopeType(pkg source.Package, def *Definition, argumentIdentifier string) (types.Type, bool) { - stdlibType := func(pkgPath, name string, pointer bool) (types.Type, bool) { - imported, ok := pkg.Import(pkgPath) - if !ok { - return nil, false - } - t := imported.Scope().Lookup(name).Type() - if pointer { - t = types.NewPointer(t) - } - return t, true - } +// defaultScopeType returns the type an argument identifier binds to: a +// reserved identifier's, as checker reports it, or string for lastEventID and +// a path value. +func defaultScopeType(checker Checker, def *Definition, argumentIdentifier string) (types.Type, bool) { switch argumentIdentifier { - case TemplateNameScopeIdentifierHTTPRequest: - return stdlibType("net/http", "Request", true) - case TemplateNameScopeIdentifierHTTPResponse: - return stdlibType("net/http", "ResponseWriter", false) - case TemplateNameScopeIdentifierContext: - return stdlibType("context", "Context", false) - case TemplateNameScopeIdentifierForm: - return stdlibType("net/url", "Values", false) - case TemplateNameScopeIdentifierMultipart: - return stdlibType("mime/multipart", "Form", true) case TemplateNameScopeIdentifierLastEventID: return types.Universe.Lookup("string").Type(), true - case TemplateNameScopeIdentifierRequestBody: - return stdlibType("io", "Reader", false) + case TemplateNameScopeIdentifierHTTPRequest, TemplateNameScopeIdentifierHTTPResponse, TemplateNameScopeIdentifierContext, + TemplateNameScopeIdentifierForm, TemplateNameScopeIdentifierMultipart, TemplateNameScopeIdentifierRequestBody: + tp, err := checker.ScopeType(argumentIdentifier) + return tp, err == nil default: - if slices.Contains(def.PathValueIdentifiers(), argumentIdentifier) { + if def.ArgumentIsPathParameter(argumentIdentifier) { return types.Universe.Lookup("string").Type(), true } return nil, false @@ -577,44 +655,40 @@ func typeQualifier(receiverPkg *types.Package) types.Qualifier { } } -func newArgumentFromIdentifier(def *Definition, pkg source.Package, arg *ast.Ident, param types.Type, qual types.Qualifier) (Argument, error) { +func newArgumentFromIdentifier(def *Definition, checker Checker, arg *ast.Ident, param types.Type, qual types.Qualifier) (Argument, error) { a := Argument{ Identifier: arg.Name, - ParamType: param, + paramType: param, } switch arg.Name { case TemplateNameScopeIdentifierContext: a.Type = ArgumentTypeRequestContext - if err := isAssignable(pkg, param, arg.Name, "context", "Context", false, qual); err != nil { + if err := bindScopeValue(&a, checker, qual); err != nil { return a, err } case TemplateNameScopeIdentifierForm: a.Type = ArgumentTypeRequestForm - bindings, err := checkFormArgument(def, pkg, param, arg.Name, "net/url", "Values", false, qual, false) - if err != nil { + if err := bindFormArgument(&a, def, checker, qual, false); err != nil { return a, err } - a.formFields = bindings case TemplateNameScopeIdentifierMultipart: a.Type = ArgumentTypeRequestMultipartForm - bindings, err := checkFormArgument(def, pkg, param, arg.Name, "mime/multipart", "Form", true, qual, true) - if err != nil { + if err := bindFormArgument(&a, def, checker, qual, true); err != nil { return a, err } - a.formFields = bindings case TemplateNameScopeIdentifierHTTPRequest: a.Type = ArgumentTypeRequest - if err := isAssignable(pkg, param, arg.Name, "net/http", "Request", true, qual); err != nil { + if err := bindScopeValue(&a, checker, qual); err != nil { return a, err } case TemplateNameScopeIdentifierHTTPResponse: a.Type = ArgumentTypeResponse - if err := isAssignable(pkg, param, arg.Name, "net/http", "ResponseWriter", false, qual); err != nil { + if err := bindScopeValue(&a, checker, qual); err != nil { return a, err } case TemplateNameScopeIdentifierLastEventID: a.Type = ArgumentTypeLastEventID - if err := checkParsedArgument(pkg, param, qual); err != nil { + if err := bindParsedArgument(&a, checker, qual); err != nil { return a, err } case TemplateNameScopeIdentifierExecute: @@ -622,13 +696,13 @@ func newArgumentFromIdentifier(def *Definition, pkg source.Package, arg *ast.Ide a.template = def.template case TemplateNameScopeIdentifierRequestBody: a.Type = ArgumentTypeRequestBody - if err := checkRequestBodyParameter(pkg, param, qual); err != nil { + if err := checkRequestBodyParameter(&a, checker, qual); err != nil { return a, err } default: - if slices.Contains(def.pathValueNames, arg.Name) { + if def.ArgumentIsPathParameter(arg.Name) { a.Type = ArgumentTypeRequestPathValue - if err := checkParsedArgument(pkg, param, qual); err != nil { + if err := bindParsedArgument(&a, checker, qual); err != nil { return a, err } return a, nil @@ -667,26 +741,18 @@ func newArgumentFromIdentifier(def *Definition, pkg source.Package, arg *ast.Ide return a, nil } -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 := imported.Scope().Lookup(name).Type() - if pointer { - t = types.NewPointer(t) - } - return t, nil -} - -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) +// bindScopeValue requires a request value -- ctx, request or response -- to +// be assignable to its parameter as it is. +func bindScopeValue(a *Argument, checker Checker, qual types.Qualifier) error { + at, err := checker.ScopeType(a.Identifier) if err != nil { return err } - if !types.AssignableTo(at, paramType) { - return fmt.Errorf("method expects type %s but %s is %s", types.TypeString(paramType, qual), argName, types.TypeString(at, qual)) + a.scopeType = at + if !types.AssignableTo(at, a.paramType) { + return fmt.Errorf("method expects type %s but %s is %s", types.TypeString(a.paramType, qual), a.Identifier, types.TypeString(at, qual)) } + a.direct = true return nil } @@ -700,7 +766,7 @@ func isSignalsCallback(def *Definition, arg *ast.Ident) bool { func (def *Definition) IsSignalsCallback(name string) bool { return def.Representation == RepresentationSSE && IsSignalsCallbackArgument(name) && - !slices.Contains(def.pathValueNames, name) + !def.ArgumentIsPathParameter(name) } // IsSignalsCallbackArgument reports whether name is a datastar patch-signals @@ -756,14 +822,16 @@ 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(pkg source.Package, param types.Type, qual types.Qualifier) error { - readerType, err := stdlibType(pkg, "io", "Reader", false) +func checkRequestBodyParameter(a *Argument, checker Checker, qual types.Qualifier) error { + readerType, err := checker.ScopeType(TemplateNameScopeIdentifierRequestBody) if err != nil { return err } - if !types.Identical(param, readerType) { - return fmt.Errorf("%s parameter must have type io.Reader, got %s", TemplateNameScopeIdentifierRequestBody, types.TypeString(param, qual)) + a.scopeType = readerType + if !types.Identical(a.paramType, readerType) { + return fmt.Errorf("%s parameter must have type io.Reader, got %s", TemplateNameScopeIdentifierRequestBody, types.TypeString(a.paramType, qual)) } + a.direct = true return nil } @@ -824,8 +892,8 @@ func checkBodyWrapperArguments(name string, call *ast.CallExpr) error { // request body, so the sugar and the explicit spelling bind identically. A // path wildcard named signals keeps its path-value meaning. It reports // whether anything was rewritten. -func rewriteSignalsArguments(call *ast.CallExpr, pathValueNames []string) bool { - if slices.Contains(pathValueNames, TemplateNameScopeIdentifierSignals) { +func rewriteSignalsArguments(call *ast.CallExpr, segments []Segment) bool { + if _, ok := pathParameter(segments, TemplateNameScopeIdentifierSignals); ok { return false } rewritten := false @@ -840,7 +908,7 @@ func rewriteSignalsArguments(call *ast.CallExpr, pathValueNames []string) bool { rewritten = true } case *ast.CallExpr: - if rewriteSignalsArguments(arg, pathValueNames) { + if rewriteSignalsArguments(arg, segments) { rewritten = true } } @@ -922,3 +990,29 @@ func patternScope() []string { TemplateNameScopeIdentifierRequestBody, } } + +// ResolveDefinitions parses the route definitions of every templates variable in pkg +// and resolves each call against receiver, or against an empty struct named +// Receiver when there is none, so handler methods are inferred. +func ResolveDefinitions(pkg source.Package, receiver *types.Named, checker Checker) ([]Definition, error) { + if receiver == nil { + receiver = asteval.NamedEmptyStruct("Receiver", pkg.Types) + } + var ( + result []Definition + errs []error + ) + for _, variable := range pkg.Variables { + defs, err := Definitions(variable) + if err != nil { + return nil, err + } + for i := range defs { + if err := ResolveCall(&defs[i], pkg, receiver, checker); err != nil { + errs = append(errs, err) + } + } + result = append(result, defs...) + } + return result, CombineErrors(errs) +} diff --git a/internal/muxt/call_internal_test.go b/internal/muxt/call_internal_test.go index 5cbb1384..aad1a88b 100644 --- a/internal/muxt/call_internal_test.go +++ b/internal/muxt/call_internal_test.go @@ -3,6 +3,8 @@ package muxt import ( "go/ast" "go/parser" + "go/token" + "go/types" "html/template" "testing" @@ -136,19 +138,19 @@ func TestPeelRepresentationWrapper(t *testing.T) { func TestRewriteSignalsArguments(t *testing.T) { for _, tt := range []struct { - expr string - pathValueNames []string - want string - rewritten bool + expr string + segments []Segment + want string + rewritten bool }{ {expr: `Save(ctx, signals)`, want: `Save(ctx, unmarshalJSON(body))`, rewritten: true}, {expr: `sse(Search(ctx, signals, sseResults))`, want: `sse(Search(ctx, unmarshalJSON(body), sseResults))`, rewritten: true}, {expr: `Save(ctx, form)`, want: `Save(ctx, form)`}, - {expr: `Show(ctx, signals)`, pathValueNames: []string{"signals"}, want: `Show(ctx, signals)`}, + {expr: `Show(ctx, signals)`, segments: []Segment{newSegment("{signals}")}, want: `Show(ctx, signals)`}, } { t.Run(tt.expr, func(t *testing.T) { call := mustParseCall(t, tt.expr) - rewritten := rewriteSignalsArguments(call, tt.pathValueNames) + rewritten := rewriteSignalsArguments(call, tt.segments) if rewritten != tt.rewritten { t.Errorf("rewriteSignalsArguments(%q) = %t, want %t", tt.expr, rewritten, tt.rewritten) } @@ -211,3 +213,23 @@ func TestDefinitionsSignalsCallback(t *testing.T) { t.Errorf("SignalsCallback() = %q, %t; want %q, true", name, ok, "countsSignals") } } + +// TestTypeQualifier states how a type is named in a message about a route: +// a type the receiver's own package declares is named on its own, and one +// from anywhere else carries its package name. +func TestTypeQualifier(t *testing.T) { + receiverPkg := types.NewPackage("example.com/server", "server") + otherPkg := types.NewPackage("example.com/other/models", "models") + qual := typeQualifier(receiverPkg) + + named := func(pkg *types.Package, name string) types.Type { + obj := types.NewTypeName(token.NoPos, pkg, name, nil) + return types.NewNamed(obj, types.NewStruct(nil, nil), nil) + } + if got := types.TypeString(named(receiverPkg, "Page"), qual); got != "Page" { + t.Errorf("a receiver package type is %q, want %q", got, "Page") + } + if got := types.TypeString(named(otherPkg, "Page"), qual); got != "models.Page" { + t.Errorf("another package's type is %q, want %q", got, "models.Page") + } +} diff --git a/internal/muxt/call_test.go b/internal/muxt/call_test.go index 9edfc99c..9a538c4f 100644 --- a/internal/muxt/call_test.go +++ b/internal/muxt/call_test.go @@ -1,58 +1,34 @@ -package muxt +package muxt_test 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/fake" + "github.com/typelate/muxt/internal/muxt" "github.com/typelate/muxt/internal/source" ) 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 := exampleTypes(t) + pkg := source.Package{Fset: fake.FileSet, Types: examplePkg} - examplePkg := packageList[0].Types - require.NotNil(t, examplePkg) - - httpPkg := findImport(examplePkg, "net/http") - require.NotNil(t, httpPkg) - httpRequestPtrType := types.NewPointer(httpPkg.Scope().Lookup("Request").Type()) - require.NotNil(t, httpRequestPtrType) - httpResponseWriterType := httpPkg.Scope().Lookup("ResponseWriter").Type() - require.NotNil(t, httpResponseWriterType) - - contextPkg := findImport(examplePkg, "context") - require.NotNil(t, contextPkg) - contextContextType := contextPkg.Scope().Lookup("Context").Type() - require.NotNil(t, contextContextType) - - netURLPkg := findImport(examplePkg, "net/url") - require.NotNil(t, contextPkg) - netURLValuesType := netURLPkg.Scope().Lookup("Values").Type() - require.NotNil(t, netURLValuesType) - - mimeMultipartPkg := findImport(examplePkg, "mime/multipart") - require.NotNil(t, mimeMultipartPkg) - multipartFormType := mimeMultipartPkg.Scope().Lookup("Form").Type() - require.NotNil(t, multipartFormType) - - ioPkg := findImport(examplePkg, "io") - require.NotNil(t, ioPkg) - ioReaderType := ioPkg.Scope().Lookup("Reader").Type() - require.NotNil(t, ioReaderType) + httpRequestPtrType := types.NewPointer(fake.Lookup(t, examplePkg, "Request")) + httpResponseWriterType := fake.Lookup(t, examplePkg, "ResponseWriter") + contextContextType := fake.Lookup(t, examplePkg, "Context") + netURLValuesType := fake.Lookup(t, examplePkg, "Values") + multipartFormType := fake.Lookup(t, examplePkg, "Form") + ioReaderType := fake.Lookup(t, examplePkg, "Reader") + // The stand-ins this package declares are what the reserved argument + // identifiers bind to, and ID is the one type that parses from text. + checker := fake.StandInChecker(t, examplePkg). + ParsesFromText(fake.Lookup(t, examplePkg, "ID")). + Fake() serverType := examplePkg.Scope().Lookup("Server").Type().(*types.Named) emptyStruct := examplePkg.Scope().Lookup("Empty").Type().(*types.Named) @@ -61,181 +37,177 @@ func TestArgument(t *testing.T) { Name string Receiver *types.Named Template string - Expect func(t *testing.T, defs []Definition, err error) + Expect func(t *testing.T, defs []muxt.Definition, err error) }{ - {Name: "no args", Receiver: serverType, Template: `{{define "GET / M()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "no args", Receiver: serverType, Template: `{{define "GET / M()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Len(t, defs[0].Arguments, 0) require.Equal(t, "M", defs[0].Identifier()) - require.NotNil(t, defs[0].sig) + require.False(t, defs[0].Signature().IsZero()) }}, - {Name: "receiver method call", Receiver: serverType, Template: `{{define "GET / M()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "receiver method call", Receiver: serverType, Template: `{{define "GET / M()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "M", defs[0].Identifier()) require.Empty(t, defs[0].Arguments) - require.True(t, defs[0].isMethod, "M is a method on the receiver") + require.True(t, defs[0].IsMethod(), "M is a method on the receiver") }}, - {Name: "package function call", Receiver: serverType, Template: `{{define "GET / FunctionContext(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "package function call", Receiver: serverType, Template: `{{define "GET / FunctionContext(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "FunctionContext", defs[0].Identifier()) - require.Equal(t, ArgumentTypeRequestContext, defs[0].Arguments[0].Type) - require.False(t, defs[0].isMethod, "FunctionContext is a package-scope function, not a receiver method") + require.Equal(t, muxt.ArgumentTypeRequestContext, defs[0].Arguments[0].Type) + require.False(t, defs[0].IsMethod(), "FunctionContext is a package-scope function, not a receiver method") }}, - {Name: "request", Receiver: serverType, Template: `{{define "GET / HTTPRequest(request)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "request", Receiver: serverType, Template: `{{define "GET / HTTPRequest(request)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) - require.NotNil(t, defs[0].sig) + require.False(t, defs[0].Signature().IsZero()) require.Len(t, defs[0].Arguments, 1) require.Equal(t, "HTTPRequest", defs[0].Identifier()) require.Equal(t, "request", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeRequest, defs[0].Arguments[0].Type) - require.True(t, types.Identical(httpRequestPtrType, defs[0].Arguments[0].ParamType)) + require.Equal(t, muxt.ArgumentTypeRequest, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(httpRequestPtrType))) }}, - {Name: "context", Receiver: serverType, Template: `{{define "GET / Context(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "context", Receiver: serverType, Template: `{{define "GET / Context(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Len(t, defs[0].Arguments, 1) require.Equal(t, "Context", defs[0].Identifier()) require.Equal(t, "ctx", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeRequestContext, defs[0].Arguments[0].Type) - require.True(t, types.Identical(contextContextType, defs[0].Arguments[0].ParamType)) + require.Equal(t, muxt.ArgumentTypeRequestContext, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(contextContextType))) }}, - {Name: "response writer", Receiver: serverType, Template: `{{define "GET / HTTPResponseWriter(response)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "response writer", Receiver: serverType, Template: `{{define "GET / HTTPResponseWriter(response)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "HTTPResponseWriter", defs[0].Identifier()) require.Equal(t, "response", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeResponse, defs[0].Arguments[0].Type) - require.True(t, types.Identical(httpResponseWriterType, defs[0].Arguments[0].ParamType)) + require.Equal(t, muxt.ArgumentTypeResponse, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(httpResponseWriterType))) }}, - {Name: "form", Receiver: serverType, Template: `{{define "GET / URLValues(form)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "form", Receiver: serverType, Template: `{{define "GET / URLValues(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "URLValues", defs[0].Identifier()) require.Equal(t, "form", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeRequestForm, defs[0].Arguments[0].Type) - require.True(t, types.Identical(netURLValuesType, defs[0].Arguments[0].ParamType)) + require.Equal(t, muxt.ArgumentTypeRequestForm, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(netURLValuesType))) }}, - {Name: "multipart value form param is parsed as a field struct and rejected", Receiver: serverType, Template: `{{define "GET / MultipartForm(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "multipart value form param is parsed as a field struct and rejected", Receiver: serverType, Template: `{{define "GET / MultipartForm(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { // A non-pointer multipart.Form falls into struct field-binding mode, // where its map fields are not parseable; raw mode requires // *multipart.Form (see "multipart raw pointer"). require.ErrorContains(t, err, "failed to generate parse statements for multipart field Value: unsupported type: map[string][]string") }}, - {Name: "multipart raw pointer", Receiver: serverType, Template: `{{define "GET / MultipartFormPtr(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "multipart raw pointer", Receiver: serverType, Template: `{{define "GET / MultipartFormPtr(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "MultipartFormPtr", defs[0].Identifier()) - require.Equal(t, ArgumentTypeRequestMultipartForm, defs[0].Arguments[0].Type) - require.True(t, types.Identical(types.NewPointer(multipartFormType), defs[0].Arguments[0].ParamType)) + require.Equal(t, muxt.ArgumentTypeRequestMultipartForm, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(types.NewPointer(multipartFormType)))) }}, - {Name: "multipart param neither struct nor pointer", Receiver: serverType, Template: `{{define "GET / String(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "multipart param neither struct nor pointer", Receiver: serverType, Template: `{{define "GET / String(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "expected multipart parameter type to be a struct") }}, - {Name: "form struct", Receiver: serverType, Template: `{{define "GET / FormStruct(form)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "form struct", Receiver: serverType, Template: `{{define "GET / FormStruct(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "FormStruct", defs[0].Identifier()) - require.Equal(t, ArgumentTypeRequestForm, defs[0].Arguments[0].Type) - require.Equal(t, "In", defs[0].Arguments[0].ParamType.(*types.Named).Obj().Name()) + require.Equal(t, muxt.ArgumentTypeRequestForm, defs[0].Arguments[0].Type) + require.Equal(t, "In", defs[0].Arguments[0].ParamType().Format(unqualified)) }}, - {Name: "form param neither struct nor url.Values", Receiver: serverType, Template: `{{define "GET / String(form)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "form param neither struct nor the form values type", Receiver: serverType, Template: `{{define "GET / String(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "expected form parameter type to be a struct") }}, - {Name: "path value", Receiver: serverType, Template: `{{define "GET /{id} String(id)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "path value", Receiver: serverType, Template: `{{define "GET /{id} String(id)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "String", defs[0].Identifier()) require.Equal(t, "id", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeRequestPathValue, defs[0].Arguments[0].Type) - basic, ok := defs[0].Arguments[0].ParamType.(*types.Basic) - require.True(t, ok) - require.Equal(t, types.String, basic.Kind()) + require.Equal(t, muxt.ArgumentTypeRequestPathValue, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().IsString()) }}, - {Name: "last event id", Receiver: serverType, Template: `{{define "GET / String(lastEventID)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "last event id", Receiver: serverType, Template: `{{define "GET / String(lastEventID)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "String", defs[0].Identifier()) require.Equal(t, "lastEventID", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeLastEventID, defs[0].Arguments[0].Type) - basic, ok := defs[0].Arguments[0].ParamType.(*types.Basic) - require.True(t, ok) - require.Equal(t, types.String, basic.Kind()) + require.Equal(t, muxt.ArgumentTypeLastEventID, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().IsString()) }}, - {Name: "body", Receiver: serverType, Template: `{{define "POST / Reader(body)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "body", Receiver: serverType, Template: `{{define "POST / Reader(body)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "Reader", defs[0].Identifier()) require.Equal(t, "body", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeRequestBody, defs[0].Arguments[0].Type) - require.True(t, types.Identical(ioReaderType, defs[0].Arguments[0].ParamType)) + require.Equal(t, muxt.ArgumentTypeRequestBody, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(ioReaderType))) }}, - {Name: "body param must be exactly io.Reader", Receiver: serverType, Template: `{{define "POST / String(body)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "body param must be exactly the reader type", Receiver: serverType, Template: `{{define "POST / String(body)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "body parameter must have type io.Reader, got string") }}, - {Name: "body on a synthesized method is io.Reader", Receiver: emptyStruct, Template: `{{define "POST / SaveBody(body)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "body on a synthesized method is the reader type", Receiver: emptyStruct, Template: `{{define "POST / SaveBody(body)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) - require.Equal(t, ArgumentTypeRequestBody, defs[0].Arguments[0].Type) - require.True(t, types.Identical(ioReaderType, defs[0].Arguments[0].ParamType)) + require.Equal(t, muxt.ArgumentTypeRequestBody, defs[0].Arguments[0].Type) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(ioReaderType))) }}, - {Name: "unmarshalJSON body decodes into the parameter type", Receiver: serverType, Template: `{{define "POST / FormStruct(unmarshalJSON(body))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "unmarshalJSON body decodes into the parameter type", Receiver: serverType, Template: `{{define "POST / FormStruct(unmarshalJSON(body))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "body", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeRequestBodyJSON, defs[0].Arguments[0].Type) - require.Equal(t, "In", defs[0].Arguments[0].ParamType.(*types.Named).Obj().Name()) + require.Equal(t, muxt.ArgumentTypeRequestBodyJSON, defs[0].Arguments[0].Type) + require.Equal(t, "In", defs[0].Arguments[0].ParamType().Format(unqualified)) }}, - {Name: "unmarshalJSON body on a synthesized method is json.RawMessage", Receiver: emptyStruct, Template: `{{define "POST / SaveJSON(unmarshalJSON(body))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "unmarshalJSON body on a synthesized method is the raw JSON type", Receiver: emptyStruct, Template: `{{define "POST / SaveJSON(unmarshalJSON(body))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) - require.Equal(t, ArgumentTypeRequestBodyJSON, defs[0].Arguments[0].Type) - require.Equal(t, "encoding/json.RawMessage", types.TypeString(defs[0].Arguments[0].ParamType, nil)) + require.Equal(t, muxt.ArgumentTypeRequestBodyJSON, defs[0].Arguments[0].Type) + require.Equal(t, "example.com.RawMessage", defs[0].Arguments[0].ParamType().Format(pathQualified)) }}, - {Name: "unmarshalForm body is the form binding", Receiver: serverType, Template: `{{define "POST / FormStruct(unmarshalForm(body))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "unmarshalForm body is the form binding", Receiver: serverType, Template: `{{define "POST / FormStruct(unmarshalForm(body))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) require.Equal(t, "form", defs[0].Arguments[0].Identifier) - require.Equal(t, ArgumentTypeRequestForm, defs[0].Arguments[0].Type) - require.Equal(t, "In", defs[0].Arguments[0].ParamType.(*types.Named).Obj().Name()) + require.Equal(t, muxt.ArgumentTypeRequestForm, defs[0].Arguments[0].Type) + require.Equal(t, "In", defs[0].Arguments[0].ParamType().Format(unqualified)) }}, - {Name: "marshalJSON with data result", Receiver: serverType, Template: `{{define "GET / marshalJSON(M())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "marshalJSON with data result", Receiver: serverType, Template: `{{define "GET / marshalJSON(M())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, RepresentationMarshalJSON, defs[0].Representation) - require.Equal(t, ResultShapeData, defs[0].ResultShape()) + require.Equal(t, muxt.RepresentationMarshalJSON, defs[0].Representation) + require.Equal(t, muxt.ResultShapeData, defs[0].ResultShape()) }}, - {Name: "marshalJSON with data and error results", Receiver: serverType, Template: `{{define "GET / marshalJSON(StringError())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "marshalJSON with data and error results", Receiver: serverType, Template: `{{define "GET / marshalJSON(StringError())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, RepresentationMarshalJSON, defs[0].Representation) - require.Equal(t, ResultShapeDataError, defs[0].ResultShape()) + require.Equal(t, muxt.RepresentationMarshalJSON, defs[0].Representation) + require.Equal(t, muxt.ResultShapeDataError, defs[0].ResultShape()) }}, - {Name: "marshalJSON requires a result", Receiver: serverType, Template: `{{define "GET / marshalJSON(NoResults())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "marshalJSON requires a result", Receiver: serverType, Template: `{{define "GET / marshalJSON(NoResults())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "marshalJSON requires a result to marshal but NoResults() returns nothing") }}, - {Name: "marshalJSON requires a non-error result", Receiver: serverType, Template: `{{define "GET / marshalJSON(NoParams())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "marshalJSON requires a non-error result", Receiver: serverType, Template: `{{define "GET / marshalJSON(NoParams())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "marshalJSON requires a non-error result but NoParams() error only returns an error") }}, - {Name: "marshalJSON second result must be an error", Receiver: serverType, Template: `{{define "GET / marshalJSON(StringOK())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "marshalJSON second result must be an error", Receiver: serverType, Template: `{{define "GET / marshalJSON(StringOK())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "marshalJSON requires the second result of StringOK() (string, bool) to be an error, got bool") }}, - {Name: "marshalJSON first result must not be an error", Receiver: serverType, Template: `{{define "GET / marshalJSON(TwoErrors())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "marshalJSON first result must not be an error", Receiver: serverType, Template: `{{define "GET / marshalJSON(TwoErrors())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "marshalJSON requires a non-error first result to marshal but TwoErrors() (error, error) returns an error value") }}, - {Name: "marshalJSON allows at most two results", Receiver: serverType, Template: `{{define "GET / marshalJSON(ThreeResults())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "marshalJSON allows at most two results", Receiver: serverType, Template: `{{define "GET / marshalJSON(ThreeResults())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "marshalJSON allows at most two results but ThreeResults() (int, int, error) has 3") }}, - {Name: "signals callback with exact error result", Receiver: serverType, Template: `{{define "GET /b sse(StreamGoodSignals(countsSignals))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "signals callback with exact error result", Receiver: serverType, Template: `{{define "GET /b sse(StreamGoodSignals(countsSignals))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ArgumentTypeSignalsCallback, defs[0].Arguments[0].Type) + require.Equal(t, muxt.ArgumentTypeSignalsCallback, defs[0].Arguments[0].Type) }}, - {Name: "signals callback result must be exactly error", Receiver: serverType, Template: `{{define "GET /b sse(StreamBadSignals(countsSignals))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "signals callback result must be exactly error", Receiver: serverType, Template: `{{define "GET /b sse(StreamBadSignals(countsSignals))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "the countsSignals signals callback must be a func(T) error") }}, - {Name: "nested method call", Receiver: serverType, Template: `{{define "GET / Any(Context(ctx))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "nested method call", Receiver: serverType, Template: `{{define "GET / Any(Context(ctx))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) // The outer Any(any) receives the result of the inner Context call. @@ -245,244 +217,297 @@ func TestArgument(t *testing.T) { require.Len(t, defs[0].Arguments, 1) require.NotEmpty(t, defs[0].Arguments[0].Identifier) - isTypeAny(t, defs[0].Arguments[0].ParamType) + isTypeAny(t, defs[0].Arguments[0].ParamType()) nested := defs[0].Arguments[0] require.Equal(t, "Context", defs[0].Arguments[0].Identifier) - require.True(t, nested.isMethod, "Context is a receiver method") - require.NotNil(t, nested.sig, "nested call signature") + require.True(t, nested.IsMethod(), "Context is a receiver method") + require.False(t, nested.Signature().IsZero(), "nested call signature") - require.Equal(t, ArgumentTypeCall, defs[0].Arguments[0].Type) + require.Equal(t, muxt.ArgumentTypeCall, defs[0].Arguments[0].Type) require.Equal(t, "Context", defs[0].Arguments[0].Identifier) - isTypeAny(t, defs[0].Arguments[0].ParamType) + isTypeAny(t, defs[0].Arguments[0].ParamType()) - require.Len(t, defs[0].Arguments[0].args, 1) - require.Equal(t, ArgumentTypeRequestContext, defs[0].Arguments[0].args[0].Type) - require.Equal(t, "ctx", defs[0].Arguments[0].args[0].Identifier) - require.Equal(t, contextContextType, defs[0].Arguments[0].args[0].ParamType) + require.Len(t, defs[0].Arguments[0].Arguments(), 1) + require.Equal(t, muxt.ArgumentTypeRequestContext, defs[0].Arguments[0].Arguments()[0].Type) + require.Equal(t, "ctx", defs[0].Arguments[0].Arguments()[0].Identifier) + require.True(t, defs[0].Arguments[0].Arguments()[0].ParamType().Identical(source.NewType(contextContextType))) }}, - {Name: "nested package function call", Receiver: serverType, Template: `{{define "GET / Any(FunctionContext(ctx))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "nested package function call", Receiver: serverType, Template: `{{define "GET / Any(FunctionContext(ctx))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) - isTypeAny(t, defs[0].Arguments[0].ParamType) + isTypeAny(t, defs[0].Arguments[0].ParamType()) - requireArgument(t, defs[0].Arguments, 0, "FunctionContext", ArgumentTypeCall, "any") + requireArgument(t, defs[0].Arguments, 0, "FunctionContext", muxt.ArgumentTypeCall, "any") nested := defs[0].Arguments[0] - require.False(t, nested.isMethod, "FunctionContext is a package-scope function") - requireArgument(t, nested.args, 0, "ctx", ArgumentTypeRequestContext, "context.Context") + require.False(t, nested.IsMethod(), "FunctionContext is a package-scope function") + requireArgument(t, nested.Arguments(), 0, "ctx", muxt.ArgumentTypeRequestContext, "example.com.Context") }}, - {Name: "synthesized method", Receiver: emptyStruct, Template: `{{define "GET / DoesNotExist(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "synthesized method", Receiver: emptyStruct, Template: `{{define "GET / DoesNotExist(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs, 1) // The receiver has no DoesNotExist method and it is not a package // function, so its signature is synthesized from the call scope. - require.NotNil(t, defs[0].sig) - require.True(t, defs[0].isMethod, "a synthesized call becomes a required receiver method") + require.False(t, defs[0].Signature().IsZero()) + require.True(t, defs[0].IsMethod(), "a synthesized call becomes a required receiver method") require.Equal(t, "ctx", defs[0].Arguments[0].Identifier) - require.True(t, types.Identical(contextContextType, defs[0].Arguments[0].ParamType)) + require.True(t, defs[0].Arguments[0].ParamType().Identical(source.NewType(contextContextType))) }}, - {Name: "synthesized method with a repeated argument", Receiver: emptyStruct, Template: `{{define "GET / RepeatedArg(request, request)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "synthesized method with a repeated argument", Receiver: emptyStruct, Template: `{{define "GET / RepeatedArg(request, request)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "cannot infer a signature for RepeatedArg: the request argument is passed more than once; define the method on the receiver to use repeated arguments") }}, - {Name: "error when argument is not assignable to parameter", Receiver: serverType, Template: `{{define "GET / Context(request)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "method expects type context.Context but request is *http.Request") + {Name: "error when argument is not assignable to parameter", Receiver: serverType, Template: `{{define "GET / Context(request)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "method expects type Context but request is *Request") }}, - {Name: "passing context argument when parameter is a string", Receiver: serverType, Template: `{{define "GET / String(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "method expects type string but ctx is context.Context") + {Name: "passing context argument when parameter is a string", Receiver: serverType, Template: `{{define "GET / String(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "method expects type string but ctx is Context") }}, - {Name: "passing request when parameter is a receiver pointer", Receiver: serverType, Template: `{{define "GET / PtrServer(request)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "method expects type *Server but request is *http.Request") + {Name: "passing request when parameter is a receiver pointer", Receiver: serverType, Template: `{{define "GET / PtrServer(request)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "method expects type *Server but request is *Request") }}, - {Name: "passing response when parameter is a string", Receiver: serverType, Template: `{{define "GET / String(response)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "method expects type string but response is http.ResponseWriter") + {Name: "passing response when parameter is a string", Receiver: serverType, Template: `{{define "GET / String(response)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "method expects type string but response is ResponseWriter") }}, - {Name: "too few arguments", Receiver: serverType, Template: `{{define "GET / String()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "too few arguments", Receiver: serverType, Template: `{{define "GET / String()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, `handler func String(string) any expects 1 arguments but call String() has 0`) require.Len(t, defs, 1) }}, - {Name: "too many arguments", Receiver: serverType, Template: `{{define "GET /{name} Context(ctx, name)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, `handler func Context(context.Context) any expects 1 arguments but call Context(ctx, name) has 2`) + {Name: "too many arguments", Receiver: serverType, Template: `{{define "GET /{name} Context(ctx, name)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, `handler func Context(Context) any expects 1 arguments but call Context(ctx, name) has 2`) }}, - {Name: "execute callback when method takes no parameters", Receiver: serverType, Template: `{{define "GET / NoParams(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "execute callback when method takes no parameters", Receiver: serverType, Template: `{{define "GET / NoParams(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "execute argument for NoParams must be a func(...) error") }}, - {Name: "wrong argument type in shared field list", Receiver: serverType, Template: `{{define "GET /post/{postID}/comment/{commentID} FieldList(ctx, request, commentID)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "method expects type string but request is *http.Request") + {Name: "wrong argument type in shared field list", Receiver: serverType, Template: `{{define "GET /post/{postID}/comment/{commentID} FieldList(ctx, request, commentID)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "method expects type string but request is *Request") }}, - {Name: "data result shape", Receiver: serverType, Template: `{{define "GET / M()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "data result shape", Receiver: serverType, Template: `{{define "GET / M()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ResultShapeData, defs[0].ResultShape()) + require.Equal(t, muxt.ResultShapeData, defs[0].ResultShape()) }}, - {Name: "data and error result shape", Receiver: serverType, Template: `{{define "GET / StringError()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "data and error result shape", Receiver: serverType, Template: `{{define "GET / StringError()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ResultShapeDataError, defs[0].ResultShape()) + require.Equal(t, muxt.ResultShapeDataError, defs[0].ResultShape()) }}, - {Name: "data and ok result shape", Receiver: serverType, Template: `{{define "GET / StringOK()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "data and ok result shape", Receiver: serverType, Template: `{{define "GET / StringOK()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ResultShapeDataOK, defs[0].ResultShape()) + require.Equal(t, muxt.ResultShapeDataOK, defs[0].ResultShape()) }}, - {Name: "method with no results", Receiver: serverType, Template: `{{define "GET / NoResults()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "method with no results", Receiver: serverType, Template: `{{define "GET / NoResults()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method NoResults() has no results; it should have one or two") }}, - {Name: "second result must be error or bool", Receiver: serverType, Template: `{{define "GET / TwoResultsSecondNotErrorOrBool()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "second result must be error or bool", Receiver: serverType, Template: `{{define "GET / TwoResultsSecondNotErrorOrBool()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "the second result of TwoResultsSecondNotErrorOrBool() (int, float64) must be an error or a bool, got float64") }}, - {Name: "execute method must return only error", Receiver: serverType, Template: `{{define "GET / ExecuteReturnsValue(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "execute method must return only error", Receiver: serverType, Template: `{{define "GET / ExecuteReturnsValue(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method ExecuteReturnsValue(func() error) (int, error) receiving the execute callback must return only error") }}, - {Name: "sse method must return nothing or an error", Receiver: serverType, Template: `{{define "GET /x sse(SSEReturnsValue(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse method must return nothing or an error", Receiver: serverType, Template: `{{define "GET /x sse(SSEReturnsValue(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "sse handler method SSEReturnsValue(func(string) error) int must return nothing or a single error") }}, - {Name: "sse method returning nothing", Receiver: serverType, Template: `{{define "GET /x sse(SSEEvents(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse method returning nothing", Receiver: serverType, Template: `{{define "GET /x sse(SSEEvents(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ResultShapeNone, defs[0].ResultShape()) + require.Equal(t, muxt.ResultShapeNone, defs[0].ResultShape()) }}, - {Name: "nested call with no results", Receiver: serverType, Template: `{{define "GET / Any(NoResults())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "nested call with no results", Receiver: serverType, Template: `{{define "GET / Any(NoResults())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method NoResults() has no results; it should have one or two") }}, - {Name: "execute callback with data parameter", Receiver: serverType, Template: `{{define "GET / ExecuteTD(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "result with a StatusCode method", Receiver: serverType, Template: `{{define "GET / Coded()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultStatusCodeMethod, defs[0].ResultStatusCode()) + }}, + {Name: "result with a StatusCode field", Receiver: serverType, Template: `{{define "GET / WithField()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultStatusCodeField, defs[0].ResultStatusCode()) + }}, + {Name: "result without a StatusCode", Receiver: serverType, Template: `{{define "GET / StringError()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultStatusCodeNone, defs[0].ResultStatusCode()) + }}, + {Name: "execute callback data with a StatusCode method", Receiver: serverType, Template: `{{define "GET / ExecuteCoded(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultStatusCodeMethod, defs[0].ResultStatusCode()) + }}, + {Name: "execute callback without data", Receiver: serverType, Template: `{{define "GET / ExecuteNoArg(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultStatusCodeNone, defs[0].ResultStatusCode()) + }}, + {Name: "nested call second result must be error or bool", Receiver: serverType, Template: `{{define "GET / Any(TwoResultsSecondNotErrorOrBool())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "the second result of TwoResultsSecondNotErrorOrBool() (int, float64) must be an error or a bool, got float64") + }}, + {Name: "execute method must not return a value in place of error", Receiver: serverType, Template: `{{define "GET / ExecuteReturnsInt(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "method ExecuteReturnsInt(func() error) int receiving the execute callback must return only error") + }}, + {Name: "nested call with data and ok results", Receiver: serverType, Template: `{{define "GET / Any(StringOK())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultShapeDataOK, defs[0].Arguments[0].ResultShape()) + }}, + {Name: "nested call with data and error results", Receiver: serverType, Template: `{{define "GET / Any(StringError())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultShapeDataError, defs[0].Arguments[0].ResultShape()) + }}, + {Name: "nested call with a data result", Receiver: serverType, Template: `{{define "GET / Any(M())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.Equal(t, muxt.ResultShapeData, defs[0].Arguments[0].ResultShape()) + }}, + {Name: "execute callback with data parameter", Receiver: serverType, Template: `{{define "GET / ExecuteTD(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.True(t, defs[0].Arguments[0].CallbackHasArg()) - named, ok := defs[0].Arguments[0].CallbackResultType().(*types.Named) - require.True(t, ok) - require.Equal(t, "TD", named.Obj().Name()) + require.Equal(t, "TD", defs[0].Arguments[0].CallbackResultType().Format(unqualified)) }}, - {Name: "execute callback without data parameter", Receiver: serverType, Template: `{{define "GET / ExecuteNoArg(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "execute callback without data parameter", Receiver: serverType, Template: `{{define "GET / ExecuteNoArg(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.False(t, defs[0].Arguments[0].CallbackHasArg()) - st, ok := defs[0].Arguments[0].CallbackResultType().(*types.Struct) - require.True(t, ok) - require.Zero(t, st.NumFields()) + require.Equal(t, "struct{}", defs[0].Arguments[0].CallbackResultType().Format(unqualified)) }}, - {Name: "execute callback parameter is not a function", Receiver: serverType, Template: `{{define "GET / ExecuteNotFunc(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "execute callback parameter is not a function", Receiver: serverType, Template: `{{define "GET / ExecuteNotFunc(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "execute argument for ExecuteNotFunc must be a func(...) error") }}, - {Name: "execute callback with too many parameters", Receiver: serverType, Template: `{{define "GET / ExecuteMultiArg(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "execute callback with too many parameters", Receiver: serverType, Template: `{{define "GET / ExecuteMultiArg(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "execute callback must have zero or one parameter; wrap multiple values in a struct") }}, - {Name: "sse callback parameter is not a function", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackNotFunc(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse callback parameter is not a function", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackNotFunc(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "execute parameter for SSECallbackNotFunc must be a function") }}, - {Name: "sse callback with too many parameters", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackMultiArg(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse callback with too many parameters", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackMultiArg(execute))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "sse callback must have zero or one parameter; wrap multiple values in a struct") }}, - {Name: "execute callback method not defined on receiver", Receiver: emptyStruct, Template: `{{define "GET / NotDefined(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "execute callback method not defined on receiver", Receiver: emptyStruct, Template: `{{define "GET / NotDefined(execute)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method NotDefined using the execute callback must be defined on the receiver type") }}, - {Name: "sse prefixed callback template missing", Receiver: serverType, Template: `{{define "GET /events sse(SSETwoCallbacks(execute, sseClock))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse prefixed callback template missing", Receiver: serverType, Template: `{{define "GET /events sse(SSETwoCallbacks(execute, sseClock))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, `no template "sseClock" for sse argument sseClock`) }}, - {Name: "sse prefixed callback template defined", Receiver: serverType, Template: `{{define "GET /events sse(SSETwoCallbacks(execute, sseClock))"}}{{end}}{{define "sseClock"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse prefixed callback template defined", Receiver: serverType, Template: `{{define "GET /events sse(SSETwoCallbacks(execute, sseClock))"}}{{end}}{{define "sseClock"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Len(t, defs[0].Arguments, 2) require.NotNil(t, defs[0].Arguments[1].Template()) require.Equal(t, "sseClock", defs[0].Arguments[1].Template().Name()) }}, - {Name: "sse message template missing", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackNotFunc(fooMessage))"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse message template missing", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackNotFunc(fooMessage))"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, `no template "fooMessage" for sse message argument fooMessage`) }}, - {Name: "sse message template defined", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackNotFunc(fooMessage))"}}{{end}}{{define "fooMessage"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "sse message template defined", Receiver: serverType, Template: `{{define "GET /x sse(SSECallbackNotFunc(fooMessage))"}}{{end}}{{define "fooMessage"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ArgumentTypeSendMessage, defs[0].Arguments[0].Type) + require.Equal(t, muxt.ArgumentTypeSendMessage, defs[0].Arguments[0].Type) require.NotNil(t, defs[0].Arguments[0].Template()) require.Equal(t, "fooMessage", defs[0].Arguments[0].Template().Name()) }}, - {Name: "method with three results", Receiver: serverType, Template: `{{define "GET / ThreeResults()"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "method with three results", Receiver: serverType, Template: `{{define "GET / ThreeResults()"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method ThreeResults() (int, int, error) has 3 results; it should have one or two") }}, - {Name: "nested call with three results", Receiver: serverType, Template: `{{define "GET / Any(ThreeResults())"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "nested call with three results", Receiver: serverType, Template: `{{define "GET / Any(ThreeResults())"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method ThreeResults() (int, int, error) has 3 results; it should have one or two") }}, - {Name: "path value with unsupported basic type", Receiver: serverType, Template: `{{define "GET /{id} Float64(id)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "path value with unsupported basic type", Receiver: serverType, Template: `{{define "GET /{id} Float64(id)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method param type float64 not supported") }}, - {Name: "path value with unsupported named type", Receiver: serverType, Template: `{{define "GET /{id} URLParam(id)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "unsupported type: url.URL") + {Name: "path value with unsupported named type", Receiver: serverType, Template: `{{define "GET /{id} URLParam(id)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "unsupported type: URL") }}, - {Name: "path value with text unmarshaler", Receiver: serverType, Template: `{{define "GET /{id} TextUnmarshalerParam(id)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "path value with text unmarshaler", Receiver: serverType, Template: `{{define "GET /{id} TextUnmarshalerParam(id)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ArgumentTypeRequestPathValue, defs[0].Arguments[0].Type) + require.Equal(t, muxt.ArgumentTypeRequestPathValue, defs[0].Arguments[0].Type) }}, - {Name: "last event id with unsupported basic type", Receiver: serverType, Template: `{{define "GET / Float64(lastEventID)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "last event id with unsupported basic type", Receiver: serverType, Template: `{{define "GET / Float64(lastEventID)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "method param type float64 not supported") }}, - {Name: "form struct with unsupported field type", Receiver: serverType, Template: `{{define "GET / FormUnsupportedField(form)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "failed to generate parse statements for form field href: unsupported type: url.URL") + {Name: "form struct with unsupported field type", Receiver: serverType, Template: `{{define "GET / FormUnsupportedField(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "failed to generate parse statements for form field href: unsupported type: URL") }}, - {Name: "multipart struct with file header fields", Receiver: serverType, Template: `{{define "POST / Upload(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "multipart struct with file header fields", Receiver: serverType, Template: `{{define "POST / Upload(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) - require.Equal(t, ArgumentTypeRequestMultipartForm, defs[0].Arguments[0].Type) + require.Equal(t, muxt.ArgumentTypeRequestMultipartForm, defs[0].Arguments[0].Type) }}, - {Name: "multipart struct with unsupported field type", Receiver: serverType, Template: `{{define "POST / BadUpload(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "failed to generate parse statements for multipart field File: unsupported type: multipart.File") + {Name: "multipart struct with unsupported field type", Receiver: serverType, Template: `{{define "POST / BadUpload(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "failed to generate parse statements for multipart field File: unsupported type: File") }}, - {Name: "form struct with file header field is unsupported", Receiver: serverType, Template: `{{define "GET / Upload(form)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { - require.ErrorContains(t, err, "failed to generate parse statements for form field File: unsupported type: *multipart.FileHeader") + {Name: "form struct with file header field is unsupported", Receiver: serverType, Template: `{{define "GET / Upload(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.ErrorContains(t, err, "failed to generate parse statements for form field File: unsupported type: *FileHeader") }}, - {Name: "form struct field bindings", Receiver: serverType, Template: `{{define "GET / TaggedForm(form)"}}{{end}}{{define "count-template"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "form struct field bindings", Receiver: serverType, Template: `{{define "GET / TaggedForm(form)"}}{{end}}{{define "count-template"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) fields := defs[0].Arguments[0].FormFields() require.Len(t, fields, 2) count := fields[0] - require.Equal(t, "Count", count.Field.Name()) + require.Equal(t, "Count", count.Name) require.Equal(t, "count-input", count.InputName) require.NotNil(t, count.Template) require.Equal(t, "count-template", count.Template.Name()) require.False(t, count.Slice) require.False(t, count.FileHeader) - require.Equal(t, UnmarshalInt, count.Method) + require.Equal(t, muxt.UnmarshalInt, count.Method) require.Len(t, count.Validations, 1) - minLength, ok := count.Validations[0].(MinLengthValidation) + minLength, ok := count.Validations[0].(muxt.MinLengthValidation) require.True(t, ok) require.Equal(t, "count-input", minLength.Name) require.Equal(t, 1, minLength.MinLength) tags := fields[1] - require.Equal(t, "Tags", tags.Field.Name()) + require.Equal(t, "Tags", tags.Name) require.Equal(t, "tag", tags.InputName) require.Nil(t, tags.Template) require.True(t, tags.Slice) - require.Equal(t, UnmarshalString, tags.Method) - basic, ok := tags.Elem.(*types.Basic) - require.True(t, ok) - require.Equal(t, types.String, basic.Kind()) + require.Equal(t, muxt.UnmarshalString, tags.Method) + require.True(t, tags.Elem().IsString()) }}, - {Name: "form field validation attribute is invalid", Receiver: serverType, Template: `{{define "GET / TaggedForm(form)"}}{{end}}{{define "count-template"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "form field validation attribute is invalid", Receiver: serverType, Template: `{{define "GET / TaggedForm(form)"}}{{end}}{{define "count-template"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.ErrorContains(t, err, "minlength must be an integer") }}, - {Name: "multipart struct field bindings", Receiver: serverType, Template: `{{define "POST / Upload(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "multipart struct field bindings", Receiver: serverType, Template: `{{define "POST / Upload(multipart)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) fields := defs[0].Arguments[0].FormFields() require.Len(t, fields, 4) - require.Equal(t, "Name", fields[0].Field.Name()) + require.Equal(t, "Name", fields[0].Name) require.False(t, fields[0].FileHeader) - require.Equal(t, "Tags", fields[1].Field.Name()) + require.Equal(t, "Tags", fields[1].Name) require.True(t, fields[1].Slice) - require.Equal(t, "File", fields[2].Field.Name()) + require.Equal(t, "File", fields[2].Name) require.True(t, fields[2].FileHeader) require.False(t, fields[2].Slice) - require.Equal(t, "Files", fields[3].Field.Name()) + require.Equal(t, "Files", fields[3].Name) require.True(t, fields[3].FileHeader) require.True(t, fields[3].Slice) }}, - {Name: "raw form param has no field bindings", Receiver: serverType, Template: `{{define "GET / URLValues(form)"}}{{end}}`, Expect: func(t *testing.T, defs []Definition, err error) { + {Name: "raw form param has no field bindings", Receiver: serverType, Template: `{{define "GET / URLValues(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { require.NoError(t, err) require.Empty(t, defs[0].Arguments[0].FormFields()) }}, + + // Direct is what generation reads to decide whether to pass a + // request value to its parameter as it is, or to parse or bind it + // first. A value the method takes as it arrives is direct. + {Name: "a request value its parameter takes as it is, is direct", Receiver: serverType, Template: `{{define "GET / Context(ctx)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.True(t, defs[0].Arguments[0].Direct(), "ctx passes to a Context parameter as it is") + }}, + {Name: "a form its parameter takes as it is, is direct", Receiver: serverType, Template: `{{define "GET / URLValues(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.True(t, defs[0].Arguments[0].Direct(), "form passes to a Values parameter as it is") + }}, + {Name: "a form bound to a struct is not direct", Receiver: serverType, Template: `{{define "GET / FormStruct(form)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.False(t, defs[0].Arguments[0].Direct(), "a struct parameter is bound field by field, not assigned") + require.NotEmpty(t, defs[0].Arguments[0].FormFields()) + }}, + {Name: "a body its parameter takes as it is, is direct", Receiver: serverType, Template: `{{define "POST / Reader(body)"}}{{end}}`, Expect: func(t *testing.T, defs []muxt.Definition, err error) { + require.NoError(t, err) + require.True(t, defs[0].Arguments[0].Direct(), "body passes to a Reader parameter as it is") + }}, } { t.Run(tc.Name, func(t *testing.T) { ts := template.Must(template.New("").Parse(tc.Template)) - defs, err := Definitions(source.Variable{Name: "templates", Set: ts}) + defs, err := muxt.Definitions(source.Variable{Name: "templates", Set: ts}) if err != nil { t.Fatal(err) } for i := range defs { - err = ResolveCall(&defs[i], source.Package{Fset: fileSet, Types: examplePkg}, tc.Receiver) + err = muxt.ResolveCall(&defs[i], pkg, tc.Receiver, checker) if err != nil { break } @@ -492,30 +517,35 @@ func TestArgument(t *testing.T) { } } -func isTypeAny(t *testing.T, tp types.Type) { +func isTypeAny(t *testing.T, tp source.Type) { t.Helper() - anyAliasType, ok := tp.(*types.Alias) - require.True(t, ok) - require.Equal(t, "any", anyAliasType.Obj().Name()) + require.Equal(t, "any", tp.Format(unqualified)) } // requireArgument asserts that the argument at index i has the expected // identifier, classification, and parameter type (compared by its type string). -func requireArgument(t *testing.T, args []Argument, i int, identifier string, argType ArgumentType, paramType string) { +func requireArgument(t *testing.T, args []muxt.Argument, i int, identifier string, argType muxt.ArgumentType, paramType string) { t.Helper() require.Greater(t, len(args), i, "argument at index %d does not exist", i) arg := args[i] - require.Equal(t, identifier, arg.Identifier, "Argument[%d].Identifier", i) - require.Equal(t, argType, arg.Type, "Argument[%d].Type", i) - require.NotNil(t, arg.ParamType, "Argument[%d].ParamType", i) - require.Equal(t, paramType, arg.ParamType.String(), "Argument[%d].ParamType", i) + require.Equal(t, identifier, arg.Identifier, "muxt.Argument[%d].Identifier", i) + require.Equal(t, argType, arg.Type, "muxt.Argument[%d].Type", i) + require.Equal(t, paramType, arg.ParamType().Format(pathQualified), "muxt.Argument[%d].ParamType", i) } -func findImport(example *types.Package, pkg string) *types.Package { - for _, p := range example.Imports() { - if p.Path() == pkg { - return p - } +// exampleTypes type checks the package in testdata/example, which declares +// its own stand-ins for the standard library types. +func exampleTypes(t *testing.T) *types.Package { + t.Helper() + files := make(map[string]string) + for _, name := range []string{"functions.go", "methods.go", "std.go"} { + src, err := os.ReadFile(filepath.Join("testdata", "example", name)) + require.NoError(t, err) + files[name] = string(src) } - return nil + return fake.Check(t, "example.com", files) } + +func unqualified(string, string) string { return "" } + +func pathQualified(_, path string) string { return path } diff --git a/internal/muxt/checker.go b/internal/muxt/checker.go new file mode 100644 index 00000000..94c2c54b --- /dev/null +++ b/internal/muxt/checker.go @@ -0,0 +1,40 @@ +package muxt + +import "go/types" + +//go:generate go run github.com/maxbrunsfeld/counterfeiter/v6 -generate + +//counterfeiter:generate --fake-name=Checker -o ../fake/checker.go . Checker + +// Checker answers what route resolution needs to know about the standard +// library: the types the reserved argument identifiers bind to, and which +// types marshal to and from text. +// +// Resolution asks these questions and nothing else of the packages outside +// the one it resolves routes in. load.StandardLibrary answers them from the +// official standard library a run loaded; a test answers them with a mock +// over types of its own, so it depends on muxt's rules rather than on the +// shape of any one standard library version. +type Checker interface { + // ScopeType returns the type a reserved argument identifier binds to: + // request (*http.Request), response (http.ResponseWriter), ctx + // (context.Context), form (url.Values), multipart (*multipart.Form) and + // body (io.Reader). + ScopeType(identifier string) (types.Type, error) + + // FileHeader returns *multipart.FileHeader, the type a multipart struct + // field binds an uploaded file to. + FileHeader() (types.Type, error) + + // RawJSON returns json.RawMessage, the parameter type inferred for + // unmarshalJSON(body) when the method is not yet defined. + RawJSON() (types.Type, error) + + // TextUnmarshaler reports whether a pointer to tp implements + // encoding.TextUnmarshaler, so tp parses from a request string. + TextUnmarshaler(tp types.Type) bool + + // TextMarshaler reports whether tp implements encoding.TextMarshaler, + // so a route path formats it as a path segment. + TextMarshaler(tp types.Type) bool +} diff --git a/internal/muxt/definition.go b/internal/muxt/definition.go index eab11d25..b8d38b58 100644 --- a/internal/muxt/definition.go +++ b/internal/muxt/definition.go @@ -146,7 +146,7 @@ func (e *ResponseWriterTemplateStateError) Error() string { if isRedirectMethod(e.Method) { remedy = "call http.Redirect in the method" } - fmt.Fprintf(&sb, "template %q calls %s but %s takes the http.ResponseWriter, so muxt writes no status code or redirect for this route: either drop the response argument or %s", + _, _ = fmt.Fprintf(&sb, "template %q calls %s but %s takes the http.ResponseWriter, so muxt writes no status code or redirect for this route: either drop the response argument or %s", e.Template, e.Method, e.Function, remedy) return sb.String() } @@ -317,7 +317,7 @@ func (e *DuplicatePatternError) MultiLineError() string { if location == "" { continue } - fmt.Fprintf(&sb, "\n%s: %s", location, note) + _, _ = fmt.Fprintf(&sb, "\n%s: %s", location, note) note = "also defined here" } return sb.String() @@ -357,11 +357,10 @@ type Definition struct { template *template.Template - pathValueTypes map[string]types.Type - pathValueNames []string - identifier string + resultStatusCode ResultStatusCode + hasResponseWriterArg bool // sourceFile is the base filename (e.g., "index.gohtml") from which this template was parsed. @@ -400,6 +399,7 @@ type Definition struct { Representation Representation + Segments []Segment Arguments []Argument } @@ -446,14 +446,51 @@ func (def Definition) DefaultStatusCode() int { return def.defaultStatus func (def Definition) MayRedirect() bool { return def.canRedirect } func (def Definition) Template() *template.Template { return def.template } func (def Definition) FunctionIdentifier() *ast.Ident { return def.fun } -func (def Definition) CallExpression() *ast.CallExpr { return def.call } +func (def Definition) CallExpression() *ast.CallExpr { return cloneCall(def.call) } func (def Definition) HasResponseWriterArg() bool { return def.hasResponseWriterArg } func (def Definition) Identifier() string { return def.identifier } func (def Definition) TemplatesVariable() string { return def.templatesVariable } -func (def Definition) Signature() *types.Signature { return def.sig } -func (def Definition) IsMethod() bool { return def.isMethod } -func (def Definition) ResultShape() ResultShape { return def.resultShape } -func (def Definition) UsesSignals() bool { return def.usesSignals } +func (def Definition) Signature() source.Type { + if def.sig == nil { + return source.Type{} + } + return source.NewType(def.sig) +} +func (def Definition) IsMethod() bool { return def.isMethod } +func (def Definition) ResultShape() ResultShape { return def.resultShape } +func (def Definition) ResultStatusCode() ResultStatusCode { return def.resultStatusCode } + +// ResultType is the type of the template data's Result field. +func (def Definition) ResultType() source.Type { return source.NewType(def.resultType()) } +func (def Definition) UsesSignals() bool { return def.usesSignals } + +func (def Definition) IsIndex() bool { + p := def.Path() + return p == "/" || p == "/{$}" +} + +// HasPathEndWildcard reports when the special path has the "{$}" wildcard +func (def Definition) HasPathEndWildcard() bool { + return strings.HasSuffix(def.Path(), "{$}") +} + +// ArgumentIsLastEventID reports whether name binds to the Last-Event-ID +// request header. The name is reserved, so no path parameter can shadow it. +func (def Definition) ArgumentIsLastEventID(name string) bool { + return name == TemplateNameScopeIdentifierLastEventID +} + +// ArgumentIsPathParameter reports whether name is a wildcard segment of the +// path. +func (def Definition) ArgumentIsPathParameter(name string) bool { + _, ok := pathParameter(def.Segments, name) + return ok +} + +// PathParameter returns the wildcard segment that names the path parameter. +func (def Definition) PathParameter(name string) (Segment, bool) { + return pathParameter(def.Segments, name) +} // SignalsCallback returns the first Signals-suffixed callback argument name, // if the route has one. @@ -461,14 +498,6 @@ func (def Definition) SignalsCallback() (string, bool) { return def.signalsCallback, def.signalsCallback != "" } -// 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 -} - // SynthesizedMethods lists the signatures ResolveCall inferred for // handler methods the receiver does not define. The results are // untyped (any), so template field checks are deferred until the @@ -506,7 +535,6 @@ func newDefinition(t *template.Template) (Definition, error, bool) { pattern: matches[templateNameMux.SubexpIndex("pattern")], fileSet: token.NewFileSet(), defaultStatusCode: http.StatusOK, - pathValueTypes: make(map[string]types.Type), template: t, spans: newNameSpans(templateNameMux.FindStringSubmatchIndex(in)), } @@ -545,23 +573,15 @@ func newDefinition(t *template.Template) (Definition, error, bool) { case "", http.MethodGet, http.MethodPost, http.MethodPut, http.MethodPatch, http.MethodDelete: } - pathValueNames := def.PathValueIdentifiers() - if err := def.checkPathValueNames(pathValueNames); err != nil { + if err := def.initializeSegments(); err != nil { return def, err, true } - def.pathValueNames = pathValueNames - err := parseHandler(def.fileSet, &def, def.pathValueNames) + err := parseHandler(def.fileSet, &def, def.Segments) if err != nil { return def, err, true } - if def.fun == nil { - for _, name := range def.pathValueNames { - def.pathValueTypes[name] = types.Universe.Lookup("string").Type() - } - } - if httpStatusCode != "" && !def.callWriteHeader(nil) { if node := findIdent(def.call, TemplateNameScopeIdentifierHTTPResponse); node != nil { return def, errAt(node, "cannot use %s as an argument and also set an HTTP status code in the template name; the handler writes the header through %[1]s", TemplateNameScopeIdentifierHTTPResponse), true @@ -572,23 +592,7 @@ func newDefinition(t *template.Template) (Definition, error, bool) { return def, nil, true } -var ( - pathSegmentPattern = regexp.MustCompile(`/\{([^}]*)}`) - templateNameMux = regexp.MustCompile(`^(?P((?P[A-Z]+)\s+)?(?P([^/])*)(?P(/(\S)*)))(\s+(?P(\d|http\.Status)\S+))?(?P.*)?$`) -) - -func (def Definition) PathValueIdentifiers() []string { - var result []string - for _, match := range pathSegmentPattern.FindAllStringSubmatch(def.path, strings.Count(def.path, "/")) { - n := match[1] - if n == "$" && strings.Count(def.path, "$") == 1 && strings.HasSuffix(def.path, "{$}") { - continue - } - n = strings.TrimSuffix(n, "...") - result = append(result, n) - } - return result -} +var templateNameMux = regexp.MustCompile(`^(?P((?P[A-Z]+)\s+)?(?P([^/])*)(?P(/(\S)*)))(\s+(?P(\d|http\.Status)\S+))?(?P.*)?$`) func hasHTTPResponseWriterArgument(call *ast.CallExpr) bool { for _, a := range call.Args { @@ -606,21 +610,6 @@ func hasHTTPResponseWriterArgument(call *ast.CallExpr) bool { return false } -func (def *Definition) checkPathValueNames(in []string) error { - for i, n := range in { - if !token.IsIdentifier(n) { - return def.pathParamErrorf(n, 0, "path parameter name not permitted: %q is not a Go identifier", n) - } - if slices.Contains(in[:i], n) { - return def.pathParamErrorf(n, 1, "path parameter name %q is used more than once; parameter names must be unique within a path", n) - } - if slices.Contains(patternScope(), n) { - return def.pathParamErrorf(n, 0, "path parameter name %s conflicts with a reserved identifier (%s)", n, strings.Join(patternScope(), ", ")) - } - } - return nil -} - func (def Definition) byPathThenMethod(d Definition) int { if n := cmp.Compare(def.path, d.path); n != 0 { return n @@ -631,7 +620,7 @@ func (def Definition) byPathThenMethod(d Definition) int { return cmp.Compare(def.handler, d.handler) } -func parseHandler(fileSet *token.FileSet, def *Definition, pathParameterNames []string) error { +func parseHandler(fileSet *token.FileSet, def *Definition, segments []Segment) error { if def.handler == "" { return nil } @@ -657,7 +646,7 @@ func parseHandler(fileSet *token.FileSet, def *Definition, pathParameterNames [] return errAt(call, "unexpected ellipsis") } - def.usesSignals = rewriteSignalsArguments(call, pathParameterNames) + def.usesSignals = rewriteSignalsArguments(call, segments) if def.Representation == RepresentationSSE { for _, a := range call.Args { if ident, ok := a.(*ast.Ident); ok && def.IsSignalsCallback(ident.Name) { @@ -667,7 +656,12 @@ func parseHandler(fileSet *token.FileSet, def *Definition, pathParameterNames [] } } - scope := append(patternScope(), pathParameterNames...) + scope := patternScope() + for _, segment := range segments { + if segment.IsWildcard() { + scope = append(scope, segment.value) + } + } slices.Sort(scope) if err := checkArguments(scope, call, def.Representation == RepresentationSSE); err != nil { return err @@ -1097,3 +1091,30 @@ func isSafeTemplateDataMethod(methodName string) bool { } return safeMethodsSet[methodName] } + +func (def Definition) ExecuteArgumentIndex() (int, bool) { + for i, arg := range def.Arguments { + if arg.Type == ArgumentTypeExecute && + arg.Identifier == TemplateNameScopeIdentifierExecute { + return i, true + } + } + return 0, false +} + +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 +} diff --git a/internal/muxt/definition_fuzz_test.go b/internal/muxt/definition_fuzz_test.go index c12b12cd..78fe9838 100644 --- a/internal/muxt/definition_fuzz_test.go +++ b/internal/muxt/definition_fuzz_test.go @@ -50,13 +50,19 @@ func FuzzNewDefinition(f *testing.F) { // intentionally not asserted here so the fuzzer keeps hunting // for panics and inconsistent error paths. _ = def.DefaultStatusCode() - // PathValueIdentifiers must not panic and must be unique. + // Wildcard segments must name unique path parameters. seen := make(map[string]struct{}) - for _, id := range def.PathValueIdentifiers() { - if _, dup := seen[id]; dup { - t.Fatalf("duplicate path value identifier %q in %q", id, name) + for _, segment := range def.Segments { + if segment.Kind() == SegmentKindUnknown { + t.Fatalf("segment of unknown kind %q in %q", segment.Value(), name) } - seen[id] = struct{}{} + if !segment.IsWildcard() { + continue + } + if _, dup := seen[segment.Value()]; dup { + t.Fatalf("duplicate path value identifier %q in %q", segment.Value(), name) + } + seen[segment.Value()] = struct{}{} } }) } diff --git a/internal/muxt/definition_internal_test.go b/internal/muxt/definition_internal_test.go index 6bce3ec9..69b8be34 100644 --- a/internal/muxt/definition_internal_test.go +++ b/internal/muxt/definition_internal_test.go @@ -698,3 +698,22 @@ func TestNewTemplateName(t *testing.T) { }) } } + +func TestDefinition_IsIndex(t *testing.T) { + t.Run("not index", func(t *testing.T) { + def := Definition{path: "/foo"} + require.False(t, def.IsIndex()) + }) + t.Run("slash", func(t *testing.T) { + def := Definition{path: "/"} + require.True(t, def.IsIndex()) + }) + t.Run("slash dollar", func(t *testing.T) { + def := Definition{path: "/{$}"} + require.True(t, def.IsIndex()) + }) + t.Run("malformed path", func(t *testing.T) { + def := Definition{path: " / "} + require.True(t, def.IsIndex()) + }) +} diff --git a/internal/muxt/resolve_test.go b/internal/muxt/resolve_test.go new file mode 100644 index 00000000..0ec3c912 --- /dev/null +++ b/internal/muxt/resolve_test.go @@ -0,0 +1,79 @@ +package muxt_test + +import ( + "go/types" + "html/template" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/fake" + "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" +) + +func TestResolveDefinitions(t *testing.T) { + const server = `package server + +type T struct{} + +type In struct{} + +func (T) Article(id int) any { return nil } +func (T) Form(In) any { return nil } +` + variable := func(name, templates string) source.Variable { + return source.Variable{Name: name, Set: template.Must(template.New(name).Parse(templates))} + } + + t.Run("every variable's routes are resolved", func(t *testing.T) { + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": server}) + src := source.Package{Fset: fake.FileSet, Types: pkg, Variables: []source.Variable{ + variable("pages", `{{define "GET /article/{id} Article(id)"}}{{end}}`), + variable("fragments", `{{define "GET /fragment/{id} Article(id)"}}{{end}}`), + }} + defs, err := muxt.ResolveDefinitions(src, fake.Lookup(t, pkg, "T").(*types.Named), fake.NewChecker().Fake()) + require.NoError(t, err) + require.Len(t, defs, 2) + require.Equal(t, "pages", defs[0].TemplatesVariable()) + require.Equal(t, "fragments", defs[1].TemplatesVariable()) + for _, def := range defs { + require.False(t, def.Signature().IsZero(), "%s is resolved", def.Name()) + segment, ok := def.PathParameter("id") + require.True(t, ok) + require.Equal(t, "int", segment.Type().Format(unqualified)) + } + }) + + t.Run("without a receiver methods are inferred", func(t *testing.T) { + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": server}) + src := source.Package{Fset: fake.FileSet, Types: pkg, Variables: []source.Variable{ + variable("templates", `{{define "GET /{id} Missing(id)"}}{{end}}`), + }} + defs, err := muxt.ResolveDefinitions(src, nil, fake.NewChecker().Fake()) + require.NoError(t, err) + require.Len(t, defs, 1) + require.Equal(t, []string{"Missing(id string) any"}, defs[0].SynthesizedMethods()) + }) + + t.Run("a route without a call is left alone", func(t *testing.T) { + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": server}) + src := source.Package{Fset: fake.FileSet, Types: pkg, Variables: []source.Variable{ + variable("templates", `{{define "GET /about"}}{{end}}`), + }} + defs, err := muxt.ResolveDefinitions(src, nil, fake.NewChecker().Fake()) + require.NoError(t, err) + require.Len(t, defs, 1) + require.True(t, defs[0].Signature().IsZero()) + }) + + t.Run("resolution errors are combined", func(t *testing.T) { + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": server}) + src := source.Package{Fset: fake.FileSet, Types: pkg, Variables: []source.Variable{ + variable("templates", `{{define "GET /a/{id} Form(id)"}}{{end}}{{define "GET /b/{name} Form(name)"}}{{end}}`), + }} + _, err := muxt.ResolveDefinitions(src, fake.Lookup(t, pkg, "T").(*types.Named), fake.NewChecker().Fake()) + require.ErrorContains(t, err, "unsupported type: In") + require.ErrorContains(t, err, "(and 1 more error)") + }) +} diff --git a/internal/muxt/segment.go b/internal/muxt/segment.go new file mode 100644 index 00000000..da7b970f --- /dev/null +++ b/internal/muxt/segment.go @@ -0,0 +1,130 @@ +package muxt + +import ( + "go/token" + "go/types" + "slices" + "strings" + + "github.com/typelate/muxt/internal/source" +) + +// SegmentKind classifies one "/"-separated part of a route pattern's path. +type SegmentKind int + +const ( + SegmentKindUnknown SegmentKind = iota + // SegmentKindLiteral matches its text exactly. + SegmentKindLiteral + // SegmentKindWildcard is a {name} segment: it matches one path segment + // and names a path parameter. + SegmentKindWildcard + // SegmentKindWildcardRemainder is a trailing {name...} segment: it + // matches the rest of the path and names a path parameter. + SegmentKindWildcardRemainder +) + +// Segment is one "/"-separated part of a route pattern's path, after any +// trailing {$}. A wildcard segment names a path parameter; once ResolveCall +// has run it also knows the type the parameter parses into. +type Segment struct { + kind SegmentKind + value string + + tp types.Type + textMarshaler bool +} + +// initializeSegments splits the path into segments and checks each wildcard +// names a distinct, unreserved Go identifier. +func (def *Definition) initializeSegments() error { + templatePath := strings.TrimSuffix(def.path, "{$}") + parts := strings.Split(templatePath, "/")[1:] + segments := make([]Segment, 0, len(parts)) + for _, part := range parts { + if part == "" { + continue + } + segment := newSegment(part) + if !segment.IsLiteral() { + if err := def.checkPathParameterName(segment, segments); err != nil { + return err + } + } + segments = append(segments, segment) + } + def.Segments = segments + return nil +} + +func (def *Definition) checkPathParameterName(segment Segment, before []Segment) error { + n := segment.value + if segment.kind == SegmentKindUnknown { + return def.pathParamErrorf(n, 0, "path segment {%s is not permitted: a wildcard is spelled {name} or {name...}", n) + } + if !token.IsIdentifier(n) { + return def.pathParamErrorf(n, 0, "path parameter name not permitted: %q is not a Go identifier", n) + } + if _, dup := pathParameter(before, n); dup { + return def.pathParamErrorf(n, 1, "path parameter name %q is used more than once; parameter names must be unique within a path", n) + } + if slices.Contains(patternScope(), n) { + return def.pathParamErrorf(n, 0, "path parameter name %s conflicts with a reserved identifier (%s)", n, strings.Join(patternScope(), ", ")) + } + return nil +} + +func newSegment(in string) Segment { + inner, ok := strings.CutPrefix(in, "{") + if !ok { + return Segment{value: in, kind: SegmentKindLiteral} + } + if value, ok := strings.CutSuffix(inner, "...}"); ok { + return Segment{value: value, kind: SegmentKindWildcardRemainder} + } + if value, ok := strings.CutSuffix(inner, "}"); ok { + return Segment{value: value, kind: SegmentKindWildcard} + } + return Segment{value: inner, kind: SegmentKindUnknown} +} + +// pathParameter returns the wildcard segment that names the path parameter. +func pathParameter(segments []Segment, name string) (Segment, bool) { + for _, segment := range segments { + if segment.IsWildcard() && segment.value == name { + return segment, true + } + } + return Segment{}, false +} + +// Value is the text of a literal segment or the parameter name of a +// wildcard segment. +func (s Segment) Value() string { return s.value } + +func (s Segment) Kind() SegmentKind { return s.kind } +func (s Segment) IsLiteral() bool { return s.kind == SegmentKindLiteral } + +// IsWildcard reports whether the segment names a path parameter, as either +// {name} or {name...}. +func (s Segment) IsWildcard() bool { + return s.kind == SegmentKindWildcard || s.kind == SegmentKindWildcardRemainder +} + +// IsRemainder reports whether the segment is a {name...} wildcard, whose +// value is the rest of the path. +func (s Segment) IsRemainder() bool { return s.kind == SegmentKindWildcardRemainder } + +// Type returns the type a wildcard segment's value parses into: the +// parameter type where the call first passes it, or string when the value +// is passed along as it arrived or is not passed to the call at all. +func (s Segment) Type() source.Type { + if s.tp == nil { + return source.NewType(types.Universe.Lookup("string").Type()) + } + return source.NewType(s.tp) +} + +// TextMarshaler reports whether Type implements encoding.TextMarshaler, so +// a route path formats the value with MarshalText. +func (s Segment) TextMarshaler() bool { return s.textMarshaler } diff --git a/internal/muxt/segment_test.go b/internal/muxt/segment_test.go new file mode 100644 index 00000000..3f434ed0 --- /dev/null +++ b/internal/muxt/segment_test.go @@ -0,0 +1,292 @@ +package muxt_test + +import ( + "fmt" + "go/types" + "html/template" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/fake" + "github.com/typelate/muxt/internal/muxt" + "github.com/typelate/muxt/internal/source" +) + +// pathValueReceiver declares methods whose parameters a path value is +// passed to, one parameter type each. +const pathValueReceiver = `package server + +// Time stands in for a type that parses from text. +type Time struct{} + +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) any { return nil } +func (T) IntString(int, string) any { return nil } +func (T) StringInt(string, int) any { return nil } +func (T) Outer(any, string) any { return nil } +func (T) Wrap(any, int) any { return nil } +func (T) Echo(string) string { return "" } +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, in which case it stays a string. +func TestPathValueTypes(t *testing.T) { + for _, tt := range []struct { + name string + definition string + param string + want string + }{ + { + name: "parsed into an int parameter", + definition: "GET /{id} Int(id)", + param: "id", + want: "int", + }, + { + name: "a string parameter needs no parsing", + definition: "GET /{id} String(id)", + param: "id", + want: "string", + }, + { + name: "a string is assignable to any", + definition: "GET /{id} Any(id)", + param: "id", + want: "string", + }, + { + name: "a text unmarshaler", + definition: "GET /{at} Time(at)", + param: "at", + want: "server.Time", + }, + { + name: "the first occurrence decides when it parses", + definition: "GET /{id} IntString(id, id)", + param: "id", + want: "int", + }, + { + name: "the first occurrence decides when it does not", + definition: "GET /{id} StringInt(id, id)", + param: "id", + want: "string", + }, + { + name: "a nested call is walked where it is passed", + definition: "GET /{id} Outer(Inner(id), id)", + param: "id", + want: "int", + }, + { + name: "a nested call that takes a string decides before a later int", + definition: "GET /{id} Wrap(Echo(id), id)", + param: "id", + want: "string", + }, + { + name: "two parameters", + definition: "GET /{a}/{b} Pair(a, b)", + param: "b", + want: "int", + }, + { + name: "a parameter the call does not pass", + definition: "GET /{a}/{b} Int(a)", + param: "b", + want: "string", + }, + { + name: "an sse-prefixed name is a callback, not a value", + definition: "GET /{sseID} Int(sseID)", + param: "sseID", + want: "string", + }, + { + name: "a remainder wildcard parses like any parameter", + definition: "GET /files/{path...} Int(path)", + param: "path", + want: "int", + }, + } { + t.Run(tt.name, func(t *testing.T) { + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": pathValueReceiver}) + receiver := pkg.Scope().Lookup("T").Type().(*types.Named) + + ts := template.Must(template.New("").Parse(fmt.Sprintf(`{{define %q}}{{end}}`, tt.definition))) + + defs, err := muxt.Definitions(source.Variable{Name: "templates", Set: ts}) + require.NoError(t, err) + require.NotEmpty(t, defs) + def := &defs[0] + require.NotNil(t, def) + + srcPkg := source.Package{Fset: fake.FileSet, Types: pkg} + fakeChecker := fake.NewChecker().ParsesFromText(fake.Lookup(t, pkg, "Time")).Fake() + + if err := muxt.ResolveCall(def, srcPkg, receiver, fakeChecker); err != nil { + t.Fatal(err) + } + + segment, ok := def.PathParameter(tt.param) + require.True(t, ok, "path parameter %q not found", tt.param) + got := segment.Type().Format(func(name, _ string) string { return name }) + require.Equal(t, tt.want, got, "wrong path parameter type") + }) + } +} + +// TestPathValueParsing states how a path value reaches its parameter: as +// the string it arrived as, or parsed by the method resolution recorded. +// Generation reads both rather than working them out again. +func TestPathValueParsing(t *testing.T) { + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": pathValueReceiver}) + receiver := pkg.Scope().Lookup("T").Type().(*types.Named) + checker := fake.NewChecker().ParsesFromText(fake.Lookup(t, pkg, "Time")).Fake() + for _, tt := range []struct { + name string + template string + wantDirect bool + wantMethod muxt.UnmarshalMethod + }{ + {name: "a string parameter takes the value as it arrived", template: "GET /{id} String(id)", wantDirect: true}, + {name: "an any parameter takes the value as it arrived", template: "GET /{id} Any(id)", wantDirect: true}, + {name: "an int parameter parses", template: "GET /{id} Int(id)", wantMethod: muxt.UnmarshalInt}, + {name: "a text unmarshaler parses", template: "GET /{at} Time(at)", wantMethod: muxt.UnmarshalTextUnmarshaler}, + } { + t.Run(tt.name, func(t *testing.T) { + 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: fake.FileSet, Types: pkg}, receiver, checker); err != nil { + t.Fatal(err) + } + argument := defs[0].Arguments[0] + if got := argument.Direct(); got != tt.wantDirect { + t.Errorf("Direct() = %t, want %t", got, tt.wantDirect) + } + if got := argument.UnmarshalMethod(); !tt.wantDirect && got != tt.wantMethod { + t.Errorf("UnmarshalMethod() = %v, want %v", got, tt.wantMethod) + } + }) + } +} + +// TestPathValueTextMarshaler states that a route path formats a path value +// with MarshalText exactly when the type it parses into is a +// TextMarshaler, as the checker says. +func TestPathValueTextMarshaler(t *testing.T) { + pkg := fake.Check(t, "example.com/server", map[string]string{"server.go": pathValueReceiver}) + receiver := pkg.Scope().Lookup("T").Type().(*types.Named) + timeType := fake.Lookup(t, pkg, "Time") + checker := fake.NewChecker().ParsesFromText(timeType).FormatsAsText(timeType).Fake() + for _, tt := range []struct { + template, param string + want bool + }{ + {template: "GET /{at} Time(at)", param: "at", want: true}, + {template: "GET /{id} Int(id)", param: "id"}, + {template: "GET /{id} String(id)", param: "id"}, + } { + t.Run(tt.template, func(t *testing.T) { + 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: fake.FileSet, Types: pkg}, receiver, checker); err != nil { + t.Fatal(err) + } + segment, ok := defs[0].PathParameter(tt.param) + if !ok { + t.Fatalf("path parameter %q not found", tt.param) + } + if got := segment.TextMarshaler(); got != tt.want { + t.Errorf("PathParameter(%q).TextMarshaler() = %t, want %t", tt.param, got, tt.want) + } + }) + } +} + +// TestSegments states how a pattern's path splits into segments and which +// wildcard spellings are rejected. +func TestSegments(t *testing.T) { + for _, tt := range []struct { + name string + definition string + want []string + wantErr string + }{ + {name: "literals and wildcards", definition: "GET /users/{id}/files/{path...}", want: []string{"literal users", "wildcard id", "literal files", "remainder path"}}, + {name: "the end wildcard is not a segment", definition: "GET /users/{$}", want: []string{"literal users"}}, + {name: "the root has no segments", definition: "GET /"}, + {name: "a literal may repeat a wildcard name", definition: "GET /id/{id}", want: []string{"literal id", "wildcard id"}}, + {name: "an unclosed wildcard", definition: "GET /{id", wantErr: "path segment {id is not permitted"}, + {name: "a wildcard followed by text", definition: "GET /{id}x", wantErr: "path segment {id}x is not permitted"}, + {name: "an empty wildcard name", definition: "GET /{}", wantErr: `"" is not a Go identifier`}, + {name: "an empty remainder name", definition: "GET /{...}", wantErr: `"" is not a Go identifier`}, + {name: "a duplicate wildcard", definition: "GET /{id}/{id}", wantErr: `path parameter name "id" is used more than once`}, + {name: "a reserved name", definition: "GET /{form}", wantErr: "path parameter name form conflicts with a reserved identifier"}, + } { + t.Run(tt.name, func(t *testing.T) { + ts := template.Must(template.New("").Parse(fmt.Sprintf(`{{define %q}}{{end}}`, tt.definition))) + defs, err := muxt.Definitions(source.Variable{Name: "templates", Set: ts}) + if tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + return + } + require.NoError(t, err) + require.Len(t, defs, 1) + var got []string + for _, segment := range defs[0].Segments { + got = append(got, describeSegment(segment)) + } + require.Equal(t, tt.want, got) + }) + } +} + +func describeSegment(segment muxt.Segment) string { + switch { + case segment.IsLiteral(): + return "literal " + segment.Value() + case segment.IsRemainder(): + return "remainder " + segment.Value() + case segment.IsWildcard(): + return "wildcard " + segment.Value() + default: + return "unknown " + segment.Value() + } +} + +func TestPathParameterLookup(t *testing.T) { + ts := template.Must(template.New("").Parse(`{{define "GET /files/{id}/{path...} M(id, path)"}}{{end}}`)) + defs, err := muxt.Definitions(source.Variable{Name: "templates", Set: ts}) + require.NoError(t, err) + def := defs[0] + + for _, name := range []string{"id", "path"} { + segment, ok := def.PathParameter(name) + require.True(t, ok, "PathParameter(%q)", name) + require.Equal(t, name, segment.Value()) + require.True(t, def.ArgumentIsPathParameter(name), "ArgumentIsPathParameter(%q)", name) + require.False(t, def.ArgumentIsLastEventID(name), "ArgumentIsLastEventID(%q)", name) + } + _, ok := def.PathParameter("files") + require.False(t, ok, "a literal segment is not a path parameter") + require.False(t, def.ArgumentIsPathParameter("files")) + require.True(t, def.ArgumentIsLastEventID(muxt.TemplateNameScopeIdentifierLastEventID)) +} diff --git a/internal/muxt/testdata/example/functions.go b/internal/muxt/testdata/example/functions.go index 98b14435..bd5333fd 100644 --- a/internal/muxt/testdata/example/functions.go +++ b/internal/muxt/testdata/example/functions.go @@ -1,20 +1,13 @@ package example -import ( - "context" - "mime/multipart" - "net/http" - "net/url" -) - -func Function() any { return nil } -func FunctionHTTPRequest(*http.Request) any { return nil } -func FunctionHTTPResponseWriter(http.ResponseWriter) any { return nil } -func FunctionContext(context.Context) any { return nil } -func FunctionString(string) any { return nil } -func FunctionAny(any) any { return nil } -func FunctionURLValues(url.Values) any { return nil } -func FunctionMultipartForm(multipart.Form) any { return nil } +func Function() any { return nil } +func FunctionHTTPRequest(*Request) any { return nil } +func FunctionHTTPResponseWriter(ResponseWriter) any { return nil } +func FunctionContext(Context) any { return nil } +func FunctionString(string) any { return nil } +func FunctionAny(any) any { return nil } +func FunctionURLValues(Values) any { return nil } +func FunctionMultipartForm(Form) any { return nil } func FunctionExecute(func() error) any { return nil } func FunctionAnyExecute(func(any) error) any { return nil } 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/testdata/example/methods.go b/internal/muxt/testdata/example/methods.go index 6dbf3b82..de586051 100644 --- a/internal/muxt/testdata/example/methods.go +++ b/internal/muxt/testdata/example/methods.go @@ -1,31 +1,21 @@ package example -import ( - "context" - "encoding" - "encoding/json" - "io" - "mime/multipart" - "net/http" - "net/url" -) - type Empty struct{} type Server struct{} -func (srv *Server) M() any { return nil } -func (srv *Server) HTTPRequest(*http.Request) any { return nil } -func (srv *Server) HTTPResponseWriter(http.ResponseWriter) any { return nil } -func (srv *Server) Context(context.Context) any { return nil } -func (srv *Server) String(string) any { return nil } -func (srv *Server) Any(any) any { return nil } -func (srv *Server) URLValues(url.Values) any { return nil } -func (srv *Server) MultipartForm(multipart.Form) any { return nil } -func (srv *Server) MultipartFormPtr(*multipart.Form) any { return nil } -func (srv *Server) PtrServer(*Server) any { return nil } -func (srv *Server) Reader(io.Reader) any { return nil } -func (srv *Server) RawJSON(json.RawMessage) any { return nil } +func (srv *Server) M() any { return nil } +func (srv *Server) HTTPRequest(*Request) any { return nil } +func (srv *Server) HTTPResponseWriter(ResponseWriter) any { return nil } +func (srv *Server) Context(Context) any { return nil } +func (srv *Server) String(string) any { return nil } +func (srv *Server) Any(any) any { return nil } +func (srv *Server) URLValues(Values) any { return nil } +func (srv *Server) MultipartForm(Form) any { return nil } +func (srv *Server) MultipartFormPtr(*Form) any { return nil } +func (srv *Server) PtrServer(*Server) any { return nil } +func (srv *Server) Reader(Reader) any { return nil } +func (srv *Server) RawJSON(RawMessage) any { return nil } // CustomError implements error to prove signals callbacks require the exact // error result type. @@ -43,13 +33,14 @@ func (srv *Server) FormStruct(In) any { return nil } func (srv *Server) NoParams() error { return nil } -func (srv *Server) FieldList(ctx context.Context, postID, commentID string) any { return nil } +func (srv *Server) FieldList(ctx Context, postID, commentID string) any { return nil } func (srv *Server) NoResults() {} func (srv *Server) TwoResultsSecondNotErrorOrBool() (int, float64) { return 0, 0 } func (srv *Server) StringOK() (string, bool) { return "", false } func (srv *Server) StringError() (string, error) { return "", nil } func (srv *Server) ExecuteReturnsValue(func() error) (int, error) { return 0, nil } +func (srv *Server) ExecuteReturnsInt(func() error) int { return 0 } func (srv *Server) SSEReturnsValue(func(string) error) int { return 0 } func (srv *Server) SSEEvents(func(string) error) {} @@ -68,33 +59,31 @@ func (srv *Server) ThreeResults() (int, int, error) { return 0, 0, nil } func (srv *Server) TwoErrors() (error, error) { return nil, nil } -func (srv *Server) Float64(float64) any { return nil } -func (srv *Server) URLParam(url.URL) any { return nil } +func (srv *Server) Float64(float64) any { return nil } +func (srv *Server) URLParam(URL) any { return nil } -// ID implements encoding.TextUnmarshaler; the interface assertion also keeps -// the encoding package in the load graph for classification. +// ID parses from text: the tests' checker says a pointer to it is a +// TextUnmarshaler. type ID [16]byte func (id *ID) UnmarshalText([]byte) error { return nil } -var _ encoding.TextUnmarshaler = (*ID)(nil) - func (srv *Server) TextUnmarshalerParam(ID) any { return nil } -type FormWithURL struct{ href url.URL } +type FormWithURL struct{ href URL } func (srv *Server) FormUnsupportedField(FormWithURL) any { return nil } type UploadForm struct { Name string Tags []string - File *multipart.FileHeader - Files []*multipart.FileHeader + File *FileHeader + Files []*FileHeader } func (srv *Server) Upload(UploadForm) any { return nil } -type BadUploadForm struct{ File multipart.File } +type BadUploadForm struct{ File File } func (srv *Server) BadUpload(BadUploadForm) any { return nil } @@ -110,3 +99,13 @@ func (srv *Server) AnyFunction(func(any) error) any { retur func (srv *Server) StringFunction(func(string) error) any { return nil } func (srv *Server) IntFunction(func(int) error) any { return nil } func (srv *Server) Functions(func(string) error, func(string) error) any { return nil } + +type Coded struct{ code int } + +func (c Coded) StatusCode() int { return c.code } + +type WithField struct{ StatusCode int } + +func (srv *Server) Coded() Coded { return Coded{} } +func (srv *Server) WithField() WithField { return WithField{} } +func (srv *Server) ExecuteCoded(func(Coded) error) error { return nil } diff --git a/internal/muxt/testdata/example/std.go b/internal/muxt/testdata/example/std.go new file mode 100644 index 00000000..45655635 --- /dev/null +++ b/internal/muxt/testdata/example/std.go @@ -0,0 +1,30 @@ +package example + +// Stand-ins for the standard library types a route argument binds to. The +// tests bind the reserved argument identifiers to these with a mock +// muxt.Checker, so what they state holds whatever the real ones look like. + +type Request struct{ Method string } + +type ResponseWriter interface{ WriteHeader(statusCode int) } + +type Context interface{ Done() <-chan struct{} } + +type Values map[string][]string + +type FileHeader struct{ Filename string } + +type Form struct { + Value map[string][]string + File map[string][]*FileHeader +} + +type File interface{ Close() error } + +type Reader interface { + Read(p []byte) (n int, err error) +} + +type RawMessage []byte + +type URL struct{ Path string } diff --git a/internal/muxt/unmarshal.go b/internal/muxt/unmarshal.go index 9c1d5592..9970c6bf 100644 --- a/internal/muxt/unmarshal.go +++ b/internal/muxt/unmarshal.go @@ -37,13 +37,11 @@ const ( UnmarshalTextUnmarshaler ) -// UnmarshalMethodFor classifies how tp parses from its string form: a basic +// 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. 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 { +// encoding.TextUnmarshaler, which checker decides. +func unmarshalMethodFor(checker Checker, tp types.Type) UnmarshalMethod { switch t := tp.(type) { case *types.Basic: switch t.Name() { @@ -77,11 +75,8 @@ func UnmarshalMethodFor(pkg source.Package, tp types.Type) UnmarshalMethod { return UnmarshalFloat64 } case *types.Named: - 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 - } + if checker.TextUnmarshaler(t) { + return UnmarshalTextUnmarshaler } } return UnmarshalUnsupported @@ -110,8 +105,8 @@ func unsupportedTypeError(tp types.Type, qual types.Qualifier, supported string) // checkUnmarshalable reports whether tp parses from a form field's // string value. -func checkUnmarshalable(pkg source.Package, tp types.Type, qual types.Qualifier) error { - if UnmarshalMethodFor(pkg, tp) != UnmarshalUnsupported { +func checkUnmarshalable(checker Checker, tp types.Type, qual types.Qualifier) error { + if unmarshalMethodFor(checker, tp) != UnmarshalUnsupported { return nil } return unsupportedTypeError(tp, qual, supportedUnmarshalFieldTypes) @@ -123,16 +118,19 @@ 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(pkg source.Package, paramType types.Type, qual types.Qualifier) error { - if isStringAssignable(paramType) { +// bindParsedArgument 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 bindParsedArgument(a *Argument, checker Checker, qual types.Qualifier) error { + a.scopeType = types.Universe.Lookup("string").Type() + if isStringAssignable(a.paramType) { + a.direct = true return nil } - switch UnmarshalMethodFor(pkg, paramType) { + a.method = unmarshalMethodFor(checker, a.paramType) + switch a.method { case UnmarshalUnsupported, UnmarshalFloat32, UnmarshalFloat64: - return unsupportedTypeError(paramType, qual, supportedUnmarshalTypes) + return unsupportedTypeError(a.paramType, qual, supportedUnmarshalTypes) default: return nil } @@ -150,16 +148,14 @@ const ( // FieldBinding describes how one struct field of a form or multipart // parameter binds to the request. type FieldBinding struct { - // Field is the bound struct field. - Field *types.Var + // Name is the bound struct field's name. + Name string // InputName is the form input name: the name struct tag or the field name. InputName string // Template is the field's validation template (template struct tag), or // nil when the tag is absent or names an undefined template. Template *template.Template - // Elem is the type parsed from one string value: the field type, or the - // slice element type when Slice is set. Undefined for FileHeader fields. - Elem types.Type + elem types.Type // Slice binds every request value for InputName, not just the first. Slice bool // FileHeader binds the field from request.MultipartForm.File instead of a @@ -172,41 +168,46 @@ type FieldBinding struct { Validations []InputValidation } -// checkFormArgument permits a form or multipart parameter to either receive +// bindFormArgument permits a form or multipart parameter to either receive // the raw request value (url.Values / *multipart.Form) or be a struct whose -// fields parse from the submitted form, returning one FieldBinding per struct -// field (nil in raw mode). Struct fields must be a supported scalar or slice +// fields parse from the submitted form, recording one FieldBinding per struct +// field (none 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, 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) +func bindFormArgument(a *Argument, def *Definition, checker Checker, qual types.Qualifier, allowFileFields bool) error { + at, err := checker.ScopeType(a.Identifier) if err != nil { - return nil, err + return err } - if types.AssignableTo(at, paramType) { - return nil, nil + a.scopeType = at + if types.AssignableTo(at, a.paramType) { + a.direct = true + return nil } - st, ok := paramType.Underlying().(*types.Struct) + st, ok := a.paramType.Underlying().(*types.Struct) if !ok { - return nil, fmt.Errorf("expected %s parameter type to be a struct", argName) + return fmt.Errorf("expected %s parameter type to be a struct", a.Identifier) } - return formStructBindings(def, pkg, st, argName, qual, allowFileFields) + bindings, err := formStructBindings(def, checker, st, a.Identifier, qual, allowFileFields) + if err != nil { + return err + } + a.formFields = bindings + return nil } -func formStructBindings(def *Definition, pkg source.Package, st *types.Struct, argName string, qual types.Qualifier, allowFileFields bool) ([]FieldBinding, error) { +func formStructBindings(def *Definition, checker Checker, st *types.Struct, argName string, qual types.Qualifier, allowFileFields bool) ([]FieldBinding, error) { var fileHeaderPtr types.Type if allowFileFields { - if mp, ok := pkg.Import("mime/multipart"); ok { - if obj := mp.Scope().Lookup("FileHeader"); obj != nil { - fileHeaderPtr = types.NewPointer(obj.Type()) - } + if fileHeader, err := checker.FileHeader(); err == nil { + fileHeaderPtr = fileHeader } } bindings := make([]FieldBinding, 0, st.NumFields()) for i := 0; i < st.NumFields(); i++ { field, tags := st.Field(i), reflect.StructTag(st.Tag(i)) fb := FieldBinding{ - Field: field, + Name: field.Name(), InputName: field.Name(), } if name, found := tags.Lookup(InputAttributeNameStructTag); found { @@ -222,20 +223,20 @@ func formStructBindings(def *Definition, pkg source.Package, st *types.Struct, a if name, found := tags.Lookup(InputAttributeTemplateStructTag); found { fb.Template = def.template.Lookup(name) } - fb.Elem = ft + fb.elem = ft if slice, ok := ft.(*types.Slice); ok { fb.Slice = true - fb.Elem = slice.Elem() + fb.elem = slice.Elem() } validations, err := fieldTemplateValidations(fb) if err != nil { return nil, err } fb.Validations = validations - if err := checkUnmarshalable(pkg, fb.Elem, qual); err != nil { + if err := checkUnmarshalable(checker, 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(pkg, fb.Elem) + fb.Method = unmarshalMethodFor(checker, fb.elem) bindings = append(bindings, fb) } return bindings, nil @@ -257,5 +258,9 @@ func fieldTemplateValidations(fb FieldBinding) ([]InputValidation, error) { if input == nil { return nil, nil } - return ParseInputValidations(fb.InputName, input, fb.Elem) + return ParseInputValidations(fb.InputName, input, fb.elem) } + +// Elem is the type parsed from one string value: the field type, or the +// slice element type when Slice is set. Undefined for FileHeader fields. +func (fb FieldBinding) Elem() source.Type { return source.NewType(fb.elem) } diff --git a/internal/source/source.go b/internal/source/source.go index 8b8781da..c161c7a4 100644 --- a/internal/source/source.go +++ b/internal/source/source.go @@ -2,12 +2,15 @@ // templates variable it declares with where its templates were written and // where they are executed. // +// It holds nothing about the standard library. What route resolution needs +// to know about that is asked of a muxt.Checker. +// // 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 -- reads a Package. It -// is plain data, with no function fields and no loader behind it, so a -// test can write one as a literal or build one from source type checked in -// memory. +// resolution, generation, type checking templates, planning mutations -- +// reads a Package. It is plain data, with no function fields 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 ( @@ -19,58 +22,19 @@ import ( // Package is a loaded Go package. type Package struct { - // Fset positions every object in Types and Imports, and every - // position in the variables. + // Fset positions every object in Types 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. diff --git a/internal/source/type.go b/internal/source/type.go new file mode 100644 index 00000000..cd42f53c --- /dev/null +++ b/internal/source/type.go @@ -0,0 +1,40 @@ +package source + +import "go/types" + +// Type is a Go type as what runs after resolution reads it: it formats to +// source text against a file's imports and answers the questions generation +// asks, so those packages need not import go/types. +type Type struct { + tp types.Type +} + +func NewType(tp types.Type) Type { return Type{tp: tp} } + +// Format writes the type as it is spelled in a file that refers to package +// path by what qualify returns for it, or unqualified when qualify returns +// "". +func (t Type) Format(qualify func(pkgName, pkgPath string) string) string { + return types.TypeString(t.tp, func(pkg *types.Package) string { + return qualify(pkg.Name(), pkg.Path()) + }) +} + +func (t Type) IsZero() bool { return t.tp == nil } + +func (t Type) Identical(u Type) bool { return types.Identical(t.tp, u.tp) } + +// IsString reports whether the underlying type is string. +func (t Type) IsString() bool { + kind, ok := t.Basic() + return ok && kind == types.String +} + +// Basic returns the kind of the underlying basic type, if it has one. +func (t Type) Basic() (types.BasicKind, bool) { + basic, ok := t.tp.Underlying().(*types.Basic) + if !ok { + return types.Invalid, false + } + return basic.Kind(), true +} diff --git a/internal/source/type_test.go b/internal/source/type_test.go new file mode 100644 index 00000000..4139f8fe --- /dev/null +++ b/internal/source/type_test.go @@ -0,0 +1,87 @@ +package source_test + +import ( + "go/token" + "go/types" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/typelate/muxt/internal/source" +) + +func named(pkg *types.Package, name string, underlying types.Type) types.Type { + return types.NewNamed(types.NewTypeName(token.NoPos, pkg, name, nil), underlying, nil) +} + +func TestTypeFormat(t *testing.T) { + out := types.NewPackage("example.com/server", "server") + model := types.NewPackage("example.com/lib/model", "model") + qualify := func(name, path string) string { + if path == out.Path() { + return "" + } + return name + "Alias" + } + errorType := types.Universe.Lookup("error").Type() + for _, tt := range []struct { + name string + tp types.Type + want string + }{ + {name: "a basic type", tp: types.Typ[types.Int], want: "int"}, + {name: "a type in the output package", tp: named(out, "T", types.NewStruct(nil, nil)), want: "T"}, + {name: "a type in another package", tp: named(model, "User", types.NewStruct(nil, nil)), want: "modelAlias.User"}, + {name: "a pointer", tp: types.NewPointer(named(model, "User", types.NewStruct(nil, nil))), want: "*modelAlias.User"}, + {name: "a slice", tp: types.NewSlice(types.Typ[types.String]), want: "[]string"}, + {name: "the empty struct", tp: types.NewStruct(nil, nil), want: "struct{}"}, + { + name: "a signature", + tp: types.NewSignatureType(nil, nil, nil, + types.NewTuple(types.NewVar(token.NoPos, nil, "", types.Typ[types.Int])), + types.NewTuple(types.NewVar(token.NoPos, nil, "", errorType)), + false), + want: "func(int) error", + }, + } { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, source.NewType(tt.tp).Format(qualify)) + }) + } +} + +func TestTypeIsString(t *testing.T) { + pkg := types.NewPackage("example.com/server", "server") + require.True(t, source.NewType(types.Typ[types.String]).IsString()) + require.True(t, source.NewType(named(pkg, "ID", types.Typ[types.String])).IsString(), "a named string type") + require.False(t, source.NewType(types.Typ[types.Int]).IsString()) + require.False(t, source.NewType(types.NewStruct(nil, nil)).IsString()) +} + +func TestTypeBasic(t *testing.T) { + pkg := types.NewPackage("example.com/server", "server") + kind, ok := source.NewType(types.Typ[types.Int8]).Basic() + require.True(t, ok) + require.Equal(t, types.Int8, kind) + + kind, ok = source.NewType(named(pkg, "Count", types.Typ[types.Uint])).Basic() + require.True(t, ok, "a named basic type") + require.Equal(t, types.Uint, kind) + + _, ok = source.NewType(types.NewStruct(nil, nil)).Basic() + require.False(t, ok) +} + +func TestTypeIdentical(t *testing.T) { + pkg := types.NewPackage("example.com/server", "server") + id := named(pkg, "ID", types.Typ[types.String]) + require.True(t, source.NewType(id).Identical(source.NewType(id))) + require.True(t, source.NewType(types.Typ[types.String]).Identical(source.NewType(types.Typ[types.String]))) + require.False(t, source.NewType(id).Identical(source.NewType(types.Typ[types.String])), "a named type is not its underlying type") + require.False(t, source.NewType(types.Typ[types.Int]).Identical(source.NewType(types.Typ[types.String]))) +} + +func TestTypeIsZero(t *testing.T) { + require.True(t, source.Type{}.IsZero()) + require.False(t, source.NewType(types.Typ[types.Int]).IsZero()) +}