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)
|
||||
}
|
||||
rec = httptest.NewRecorder()
|
||||
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/12", nil))
|
||||
if rec.Code != http.StatusOK || rec.Body.String() != `{"id":12}` {
|
||||
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/99", nil))
|
||||
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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -59,13 +59,19 @@ func (p *Plugin) Boot(app *backpack.App) error {
|
||||
func (p *Plugin) Routes(r pact.Router) error {
|
||||
r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) {
|
||||
id, ok := surf.IntParam(req, "id")
|
||||
if !ok {
|
||||
if !ok || id != 1 {
|
||||
http.NotFound(w, req)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
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
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ package lagoon
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"git.golem15.com/golem15/summercms/backpack"
|
||||
"git.golem15.com/golem15/summercms/bonfire"
|
||||
@@ -61,7 +62,7 @@ func RuntimeCommands(app *backpack.App, plugins []party.Plugin) []bonfire.Comman
|
||||
for _, row := range rows {
|
||||
ids := "(none)"
|
||||
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})
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode"
|
||||
@@ -13,6 +14,13 @@ import (
|
||||
|
||||
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.
|
||||
func HistoryTableName(pluginID string) (string, error) {
|
||||
if err := validatePluginID(pluginID); err != nil {
|
||||
@@ -83,17 +91,23 @@ func RollbackLast(gdb *gorm.DB, plugins []party.Plugin, pluginID string) error {
|
||||
pluginID = lastMigrationPlugin(plugins)
|
||||
}
|
||||
if pluginID == "" {
|
||||
return fmt.Errorf("lagoon: no plugin migrations to roll back")
|
||||
return ErrNoMigrations
|
||||
}
|
||||
var target party.Plugin
|
||||
for _, p := range plugins {
|
||||
if p.ID() != pluginID {
|
||||
continue
|
||||
if p.ID() == pluginID {
|
||||
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 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
|
||||
}
|
||||
@@ -101,8 +115,6 @@ func RollbackLast(gdb *gorm.DB, plugins []party.Plugin, pluginID string) error {
|
||||
return fmt.Errorf("lagoon: rollback %s: %w", pluginID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("lagoon: plugin %q is not activated", pluginID)
|
||||
}
|
||||
|
||||
func lastMigrationPlugin(plugins []party.Plugin) string {
|
||||
|
||||
@@ -2,12 +2,16 @@ package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.golem15.com/golem15/summercms/backpack"
|
||||
"git.golem15.com/golem15/summercms/bonfire"
|
||||
"git.golem15.com/golem15/summercms/party"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestHistoryTableName(t *testing.T) {
|
||||
@@ -95,3 +99,27 @@ func TestOpenRequiresDSN(t *testing.T) {
|
||||
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 {
|
||||
Group(prefix string, middleware []string, fn func(Router))
|
||||
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.
|
||||
|
||||
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 (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"git.golem15.com/golem15/summercms/backpack"
|
||||
@@ -28,6 +27,7 @@ type route struct {
|
||||
path string
|
||||
handler http.Handler
|
||||
middleware []string
|
||||
constraints []Constraint
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
|
||||
// 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) {
|
||||
full := joinPath(prefix, path)
|
||||
key := "GET " + full
|
||||
@@ -163,7 +219,7 @@ func (r *Router) compile() (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 = orgSlot(h)
|
||||
for i := len(rt.middleware) - 1; i >= 0; i-- {
|
||||
@@ -290,23 +346,6 @@ func noOpLimit(next http.Handler) http.Handler {
|
||||
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 {
|
||||
prefix = strings.TrimSuffix(prefix, "/")
|
||||
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) {
|
||||
r := New(nil)
|
||||
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)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
r.Where("id", `[0-9]+`)
|
||||
h, err := r.compile()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -138,6 +127,56 @@ func TestTypedIDRouteReturns404(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/nope", nil))
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
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