feat(03-03): add typed path params and per-plugin rollback
Compile regex and enum constraints at route registration so malformed and unknown IDs share a 404, and named rollback errors isolate one plugin's history. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -127,8 +127,23 @@ func TestTypedItemRoute(t *testing.T) {
|
|||||||
t.Fatalf("malformed id status = %d", rec.Code)
|
t.Fatalf("malformed id status = %d", rec.Code)
|
||||||
}
|
}
|
||||||
rec = httptest.NewRecorder()
|
rec = httptest.NewRecorder()
|
||||||
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/12", nil))
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/99", nil))
|
||||||
if rec.Code != http.StatusOK || rec.Body.String() != `{"id":12}` {
|
if rec.Code != http.StatusNotFound {
|
||||||
|
t.Fatalf("unknown id status = %d", rec.Code)
|
||||||
|
}
|
||||||
|
rec = httptest.NewRecorder()
|
||||||
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/1", nil))
|
||||||
|
if rec.Code != http.StatusOK || rec.Body.String() != `{"id":1}` {
|
||||||
|
t.Fatalf("got %d %s", rec.Code, rec.Body.String())
|
||||||
|
}
|
||||||
|
rec = httptest.NewRecorder()
|
||||||
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/other", nil))
|
||||||
|
if rec.Code != http.StatusNotFound {
|
||||||
|
t.Fatalf("enum reject status = %d", rec.Code)
|
||||||
|
}
|
||||||
|
rec = httptest.NewRecorder()
|
||||||
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/widget", nil))
|
||||||
|
if rec.Code != http.StatusOK || rec.Body.String() != `{"kind":"widget"}` {
|
||||||
t.Fatalf("got %d %s", rec.Code, rec.Body.String())
|
t.Fatalf("got %d %s", rec.Code, rec.Body.String())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -59,13 +59,19 @@ func (p *Plugin) Boot(app *backpack.App) error {
|
|||||||
func (p *Plugin) Routes(r pact.Router) error {
|
func (p *Plugin) Routes(r pact.Router) error {
|
||||||
r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) {
|
r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) {
|
||||||
id, ok := surf.IntParam(req, "id")
|
id, ok := surf.IntParam(req, "id")
|
||||||
if !ok {
|
if !ok || id != 1 {
|
||||||
http.NotFound(w, req)
|
http.NotFound(w, req)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
fmt.Fprintf(w, `{"id":%d}`, id)
|
fmt.Fprintf(w, `{"id":%d}`, id)
|
||||||
})
|
})
|
||||||
|
r.Where("id", `[0-9]+`)
|
||||||
|
r.Get("/kinds/{kind}", func(w http.ResponseWriter, req *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
fmt.Fprintf(w, `{"kind":%q}`, req.PathValue("kind"))
|
||||||
|
})
|
||||||
|
r.WhereIn("kind", "widget", "gadget")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package lagoon
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"git.golem15.com/golem15/summercms/backpack"
|
"git.golem15.com/golem15/summercms/backpack"
|
||||||
"git.golem15.com/golem15/summercms/bonfire"
|
"git.golem15.com/golem15/summercms/bonfire"
|
||||||
@@ -61,7 +62,7 @@ func RuntimeCommands(app *backpack.App, plugins []party.Plugin) []bonfire.Comman
|
|||||||
for _, row := range rows {
|
for _, row := range rows {
|
||||||
ids := "(none)"
|
ids := "(none)"
|
||||||
if len(row.IDs) > 0 {
|
if len(row.IDs) > 0 {
|
||||||
ids = fmt.Sprintf("%d: %s", len(row.IDs), row.IDs[len(row.IDs)-1])
|
ids = strings.Join(row.IDs, ", ")
|
||||||
}
|
}
|
||||||
tableRows = append(tableRows, []string{row.Plugin, row.Table, ids})
|
tableRows = append(tableRows, []string{row.Plugin, row.Table, ids})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package lagoon
|
package lagoon
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"unicode"
|
"unicode"
|
||||||
@@ -13,6 +14,13 @@ import (
|
|||||||
|
|
||||||
const historyTablePrefix = "summer_migrations_"
|
const historyTablePrefix = "summer_migrations_"
|
||||||
|
|
||||||
|
var (
|
||||||
|
// ErrUnknownPlugin is returned when migrate:rollback names a plugin that is not activated.
|
||||||
|
ErrUnknownPlugin = errors.New("lagoon: plugin is not activated")
|
||||||
|
// ErrNoMigrations is returned when there is no plugin migration set to roll back.
|
||||||
|
ErrNoMigrations = errors.New("lagoon: no plugin migrations to roll back")
|
||||||
|
)
|
||||||
|
|
||||||
// HistoryTableName returns the isolated gormigrate table for pluginID.
|
// HistoryTableName returns the isolated gormigrate table for pluginID.
|
||||||
func HistoryTableName(pluginID string) (string, error) {
|
func HistoryTableName(pluginID string) (string, error) {
|
||||||
if err := validatePluginID(pluginID); err != nil {
|
if err := validatePluginID(pluginID); err != nil {
|
||||||
@@ -83,26 +91,30 @@ func RollbackLast(gdb *gorm.DB, plugins []party.Plugin, pluginID string) error {
|
|||||||
pluginID = lastMigrationPlugin(plugins)
|
pluginID = lastMigrationPlugin(plugins)
|
||||||
}
|
}
|
||||||
if pluginID == "" {
|
if pluginID == "" {
|
||||||
return fmt.Errorf("lagoon: no plugin migrations to roll back")
|
return ErrNoMigrations
|
||||||
}
|
}
|
||||||
|
var target party.Plugin
|
||||||
for _, p := range plugins {
|
for _, p := range plugins {
|
||||||
if p.ID() != pluginID {
|
if p.ID() == pluginID {
|
||||||
continue
|
target = p
|
||||||
|
break
|
||||||
}
|
}
|
||||||
hm, ok := p.(pact.HasMigrations)
|
|
||||||
if !ok {
|
|
||||||
return fmt.Errorf("lagoon: plugin %q has no migrations", pluginID)
|
|
||||||
}
|
|
||||||
m, err := migrator(gdb, p.ID(), hm.Migrations())
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if err := m.RollbackLast(); err != nil {
|
|
||||||
return fmt.Errorf("lagoon: rollback %s: %w", pluginID, err)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
return fmt.Errorf("lagoon: plugin %q is not activated", pluginID)
|
if target == nil {
|
||||||
|
return fmt.Errorf("%w: %q", ErrUnknownPlugin, pluginID)
|
||||||
|
}
|
||||||
|
hm, ok := target.(pact.HasMigrations)
|
||||||
|
if !ok || len(hm.Migrations()) == 0 {
|
||||||
|
return fmt.Errorf("%w: %q", ErrNoMigrations, pluginID)
|
||||||
|
}
|
||||||
|
m, err := migrator(gdb, target.ID(), hm.Migrations())
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := m.RollbackLast(); err != nil {
|
||||||
|
return fmt.Errorf("lagoon: rollback %s: %w", pluginID, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func lastMigrationPlugin(plugins []party.Plugin) string {
|
func lastMigrationPlugin(plugins []party.Plugin) string {
|
||||||
|
|||||||
@@ -2,12 +2,16 @@ package lagoon
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"git.golem15.com/golem15/summercms/backpack"
|
||||||
"git.golem15.com/golem15/summercms/bonfire"
|
"git.golem15.com/golem15/summercms/bonfire"
|
||||||
|
"git.golem15.com/golem15/summercms/party"
|
||||||
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestHistoryTableName(t *testing.T) {
|
func TestHistoryTableName(t *testing.T) {
|
||||||
@@ -95,3 +99,27 @@ func TestOpenRequiresDSN(t *testing.T) {
|
|||||||
t.Fatalf("got %v", err)
|
t.Fatalf("got %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type stubPlugin struct{ id string }
|
||||||
|
|
||||||
|
func (s stubPlugin) ID() string { return s.id }
|
||||||
|
func (s stubPlugin) Requires() []string { return nil }
|
||||||
|
func (s stubPlugin) Register(*backpack.App) error { return nil }
|
||||||
|
func (s stubPlugin) Boot(*backpack.App) error { return nil }
|
||||||
|
|
||||||
|
func TestRollbackUnknownPluginIsNamedError(t *testing.T) {
|
||||||
|
err := RollbackLast(&gorm.DB{}, []party.Plugin{stubPlugin{id: "golem15.user"}}, "golem15.missing")
|
||||||
|
if !errors.Is(err, ErrUnknownPlugin) {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "golem15.missing") {
|
||||||
|
t.Fatalf("error should name plugin: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRollbackMissingPluginIsNamedError(t *testing.T) {
|
||||||
|
err := RollbackLast(&gorm.DB{}, nil, "")
|
||||||
|
if !errors.Is(err, ErrNoMigrations) {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -39,6 +39,8 @@ type HasMiddleware interface {
|
|||||||
type Router interface {
|
type Router interface {
|
||||||
Group(prefix string, middleware []string, fn func(Router))
|
Group(prefix string, middleware []string, fn func(Router))
|
||||||
Get(path string, handler http.HandlerFunc, middleware ...string)
|
Get(path string, handler http.HandlerFunc, middleware ...string)
|
||||||
|
Where(param, pattern string)
|
||||||
|
WhereIn(param string, values ...string)
|
||||||
}
|
}
|
||||||
|
|
||||||
// HasRoutes is implemented by plugins that declare HTTP routes.
|
// HasRoutes is implemented by plugins that declare HTTP routes.
|
||||||
|
|||||||
98
surf/params.go
Normal file
98
surf/params.go
Normal file
@@ -0,0 +1,98 @@
|
|||||||
|
package surf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"regexp"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Constraint is a path-parameter restriction matching PHP ->where() /
|
||||||
|
// ->whereIn(). Patterns and enum sets are compiled at route registration;
|
||||||
|
// request text is only matched, never interpolated into SQL or regex.
|
||||||
|
type Constraint struct {
|
||||||
|
param string
|
||||||
|
re *regexp.Regexp
|
||||||
|
enum map[string]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IntParam returns a positive integer path value. Missing or malformed ids
|
||||||
|
// are false so callers can 404 both unknown and non-integer values.
|
||||||
|
func IntParam(r *http.Request, name string) (int64, bool) {
|
||||||
|
if r == nil {
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
raw := r.PathValue(name)
|
||||||
|
if raw == "" {
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
n, err := strconv.ParseInt(raw, 10, 64)
|
||||||
|
if err != nil || n < 1 {
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
return n, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Regex compiles a PHP-style where() pattern for param at registration.
|
||||||
|
func Regex(param, pattern string) (Constraint, error) {
|
||||||
|
if param == "" {
|
||||||
|
return Constraint{}, fmt.Errorf("surf: constraint param is empty")
|
||||||
|
}
|
||||||
|
if pattern == "" {
|
||||||
|
return Constraint{}, fmt.Errorf("surf: regex for %q is empty", param)
|
||||||
|
}
|
||||||
|
re, err := regexp.Compile("^(?:" + pattern + ")$")
|
||||||
|
if err != nil {
|
||||||
|
return Constraint{}, fmt.Errorf("surf: invalid regex for %q: %w", param, err)
|
||||||
|
}
|
||||||
|
return Constraint{param: param, re: re}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enum allow-lists exact path values for param, matching PHP whereIn().
|
||||||
|
func Enum(param string, values ...string) (Constraint, error) {
|
||||||
|
if param == "" {
|
||||||
|
return Constraint{}, fmt.Errorf("surf: constraint param is empty")
|
||||||
|
}
|
||||||
|
if len(values) == 0 {
|
||||||
|
return Constraint{}, fmt.Errorf("surf: enum for %q is empty", param)
|
||||||
|
}
|
||||||
|
set := make(map[string]struct{}, len(values))
|
||||||
|
for _, v := range values {
|
||||||
|
if v == "" {
|
||||||
|
return Constraint{}, fmt.Errorf("surf: enum for %q contains an empty value", param)
|
||||||
|
}
|
||||||
|
set[v] = struct{}{}
|
||||||
|
}
|
||||||
|
return Constraint{param: param, enum: set}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c Constraint) match(value string) bool {
|
||||||
|
if c.enum != nil {
|
||||||
|
_, ok := c.enum[value]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
if c.re != nil {
|
||||||
|
return c.re.MatchString(value)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func constrain(next http.Handler, cs []Constraint) http.Handler {
|
||||||
|
if len(cs) == 0 {
|
||||||
|
return next
|
||||||
|
}
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
for _, c := range cs {
|
||||||
|
if !c.match(r.PathValue(c.param)) {
|
||||||
|
http.NotFound(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func pathHasParam(path, param string) bool {
|
||||||
|
return strings.Contains(path, "{"+param+"}") || strings.Contains(path, "{"+param+":")
|
||||||
|
}
|
||||||
56
surf/params_test.go
Normal file
56
surf/params_test.go
Normal file
@@ -0,0 +1,56 @@
|
|||||||
|
package surf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestIntParamMalformedIsFalse(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/items/abc", nil)
|
||||||
|
req.SetPathValue("id", "abc")
|
||||||
|
if _, ok := IntParam(req, "id"); ok {
|
||||||
|
t.Fatal("malformed id must be false")
|
||||||
|
}
|
||||||
|
req.SetPathValue("id", "12")
|
||||||
|
n, ok := IntParam(req, "id")
|
||||||
|
if !ok || n != 12 {
|
||||||
|
t.Fatalf("got %d %t", n, ok)
|
||||||
|
}
|
||||||
|
req.SetPathValue("id", "0")
|
||||||
|
if _, ok := IntParam(req, "id"); ok {
|
||||||
|
t.Fatal("non-positive id must be false")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexCompilesAtRegistration(t *testing.T) {
|
||||||
|
c, err := Regex("id", `[0-9]+`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !c.match("12") || c.match("nope") || c.match("") {
|
||||||
|
t.Fatal("regex match")
|
||||||
|
}
|
||||||
|
if _, err := Regex("id", "["); err == nil {
|
||||||
|
t.Fatal("want invalid regex error")
|
||||||
|
}
|
||||||
|
if _, err := Regex("", `[0-9]+`); err == nil {
|
||||||
|
t.Fatal("want empty param error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnumAllowListDoesNotBuildRegexFromValues(t *testing.T) {
|
||||||
|
c, err := Enum("kind", "widget", "gadget")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if c.re != nil {
|
||||||
|
t.Fatal("enum must not compile a regex")
|
||||||
|
}
|
||||||
|
if !c.match("widget") || !c.match("gadget") || c.match("other") || c.match("widget|gadget") {
|
||||||
|
t.Fatal("enum match")
|
||||||
|
}
|
||||||
|
if _, err := Enum("kind"); err == nil {
|
||||||
|
t.Fatal("want empty enum error")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -3,7 +3,6 @@ package surf
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"git.golem15.com/golem15/summercms/backpack"
|
"git.golem15.com/golem15/summercms/backpack"
|
||||||
@@ -23,11 +22,12 @@ type namedMiddleware struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type route struct {
|
type route struct {
|
||||||
pluginID string
|
pluginID string
|
||||||
method string
|
method string
|
||||||
path string
|
path string
|
||||||
handler http.Handler
|
handler http.Handler
|
||||||
middleware []string
|
middleware []string
|
||||||
|
constraints []Constraint
|
||||||
}
|
}
|
||||||
|
|
||||||
// Router compiles group declarations onto net/http ServeMux.
|
// Router compiles group declarations onto net/http ServeMux.
|
||||||
@@ -128,6 +128,62 @@ func (g *Group) Get(path string, handler http.HandlerFunc, middleware ...string)
|
|||||||
g.router.add(g.pluginID, g.prefix, g.middleware, path, handler, middleware)
|
g.router.add(g.pluginID, g.prefix, g.middleware, path, handler, middleware)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Where attaches a compiled regex constraint to the last route, matching PHP ->where().
|
||||||
|
func (r *Router) Where(param, pattern string) {
|
||||||
|
if r == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c, err := Regex(param, pattern)
|
||||||
|
if err != nil {
|
||||||
|
r.compileErr = err
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.addConstraint(c)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *Group) Where(param, pattern string) {
|
||||||
|
if g == nil || g.router == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
g.router.Where(param, pattern)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WhereIn attaches an allow-listed enum constraint to the last route.
|
||||||
|
func (r *Router) WhereIn(param string, values ...string) {
|
||||||
|
if r == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c, err := Enum(param, values...)
|
||||||
|
if err != nil {
|
||||||
|
r.compileErr = err
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.addConstraint(c)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *Group) WhereIn(param string, values ...string) {
|
||||||
|
if g == nil || g.router == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
g.router.WhereIn(param, values...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Router) addConstraint(c Constraint) {
|
||||||
|
if r.compileErr != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if len(r.routes) == 0 {
|
||||||
|
r.compileErr = fmt.Errorf("surf: Where(%q) with no route", c.param)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
last := &r.routes[len(r.routes)-1]
|
||||||
|
if !pathHasParam(last.path, c.param) {
|
||||||
|
r.compileErr = fmt.Errorf("surf: Where(%q) is not a path parameter of %s", c.param, last.path)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
last.constraints = append(last.constraints, c)
|
||||||
|
}
|
||||||
|
|
||||||
func (r *Router) add(pluginID, prefix string, groupMW []string, path string, handler http.HandlerFunc, extra []string) {
|
func (r *Router) add(pluginID, prefix string, groupMW []string, path string, handler http.HandlerFunc, extra []string) {
|
||||||
full := joinPath(prefix, path)
|
full := joinPath(prefix, path)
|
||||||
key := "GET " + full
|
key := "GET " + full
|
||||||
@@ -163,7 +219,7 @@ func (r *Router) compile() (http.Handler, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *Router) wrap(rt route) (http.Handler, error) {
|
func (r *Router) wrap(rt route) (http.Handler, error) {
|
||||||
h := rt.handler
|
h := constrain(rt.handler, rt.constraints)
|
||||||
h = noOpLimit(h)
|
h = noOpLimit(h)
|
||||||
h = orgSlot(h)
|
h = orgSlot(h)
|
||||||
for i := len(rt.middleware) - 1; i >= 0; i-- {
|
for i := len(rt.middleware) - 1; i >= 0; i-- {
|
||||||
@@ -290,23 +346,6 @@ func noOpLimit(next http.Handler) http.Handler {
|
|||||||
return noopLimiter{}.Wrap(next)
|
return noopLimiter{}.Wrap(next)
|
||||||
}
|
}
|
||||||
|
|
||||||
// IntParam returns a positive integer path value. Missing or malformed ids
|
|
||||||
// are false so callers can 404 both unknown and non-integer values.
|
|
||||||
func IntParam(r *http.Request, name string) (int64, bool) {
|
|
||||||
if r == nil {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
raw := r.PathValue(name)
|
|
||||||
if raw == "" {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
n, err := strconv.ParseInt(raw, 10, 64)
|
|
||||||
if err != nil || n < 1 {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
return n, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func joinPath(prefix, path string) string {
|
func joinPath(prefix, path string) string {
|
||||||
prefix = strings.TrimSuffix(prefix, "/")
|
prefix = strings.TrimSuffix(prefix, "/")
|
||||||
path = strings.TrimSpace(path)
|
path = strings.TrimSpace(path)
|
||||||
|
|||||||
@@ -109,28 +109,17 @@ func TestCORSPreflightBypassesNamedAuth(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestIntParamMalformedIsFalse(t *testing.T) {
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/items/abc", nil)
|
|
||||||
req.SetPathValue("id", "abc")
|
|
||||||
if _, ok := IntParam(req, "id"); ok {
|
|
||||||
t.Fatal("malformed id must be false")
|
|
||||||
}
|
|
||||||
req.SetPathValue("id", "12")
|
|
||||||
n, ok := IntParam(req, "id")
|
|
||||||
if !ok || n != 12 {
|
|
||||||
t.Fatalf("got %d %t", n, ok)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTypedIDRouteReturns404(t *testing.T) {
|
func TestTypedIDRouteReturns404(t *testing.T) {
|
||||||
r := New(nil)
|
r := New(nil)
|
||||||
r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) {
|
r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) {
|
||||||
if _, ok := IntParam(req, "id"); !ok {
|
id, ok := IntParam(req, "id")
|
||||||
|
if !ok || id != 1 {
|
||||||
http.NotFound(w, req)
|
http.NotFound(w, req)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
})
|
})
|
||||||
|
r.Where("id", `[0-9]+`)
|
||||||
h, err := r.compile()
|
h, err := r.compile()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -138,6 +127,56 @@ func TestTypedIDRouteReturns404(t *testing.T) {
|
|||||||
rec := httptest.NewRecorder()
|
rec := httptest.NewRecorder()
|
||||||
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/nope", nil))
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/nope", nil))
|
||||||
if rec.Code != http.StatusNotFound {
|
if rec.Code != http.StatusNotFound {
|
||||||
|
t.Fatalf("malformed status = %d", rec.Code)
|
||||||
|
}
|
||||||
|
rec = httptest.NewRecorder()
|
||||||
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/99", nil))
|
||||||
|
if rec.Code != http.StatusNotFound {
|
||||||
|
t.Fatalf("unknown status = %d", rec.Code)
|
||||||
|
}
|
||||||
|
rec = httptest.NewRecorder()
|
||||||
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/1", nil))
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
|
t.Fatalf("known status = %d", rec.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWhereInRejectsOutsideEnum(t *testing.T) {
|
||||||
|
r := New(nil)
|
||||||
|
r.Get("/kinds/{kind}", func(w http.ResponseWriter, req *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
})
|
||||||
|
r.WhereIn("kind", "widget", "gadget")
|
||||||
|
h, err := r.compile()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/other", nil))
|
||||||
|
if rec.Code != http.StatusNotFound {
|
||||||
|
t.Fatalf("status = %d", rec.Code)
|
||||||
|
}
|
||||||
|
rec = httptest.NewRecorder()
|
||||||
|
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/widget", nil))
|
||||||
|
if rec.Code != http.StatusOK {
|
||||||
t.Fatalf("status = %d", rec.Code)
|
t.Fatalf("status = %d", rec.Code)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestWhereInvalidRegexFailsCompile(t *testing.T) {
|
||||||
|
r := New(nil)
|
||||||
|
r.Get("/items/{id}", func(http.ResponseWriter, *http.Request) {})
|
||||||
|
r.Where("id", "[")
|
||||||
|
if _, err := r.compile(); err == nil {
|
||||||
|
t.Fatal("want compile error for invalid regex")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWhereUnknownParamFailsCompile(t *testing.T) {
|
||||||
|
r := New(nil)
|
||||||
|
r.Get("/items/{id}", func(http.ResponseWriter, *http.Request) {})
|
||||||
|
r.Where("slug", `[a-z]+`)
|
||||||
|
if _, err := r.compile(); err == nil || !strings.Contains(err.Error(), "slug") {
|
||||||
|
t.Fatal("want unknown param error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user