diff --git a/examples/hello/hello_test.go b/examples/hello/hello_test.go index 1c7b899..bac8255 100644 --- a/examples/hello/hello_test.go +++ b/examples/hello/hello_test.go @@ -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()) } } diff --git a/examples/hello/plugins/greeter/plugin.go b/examples/hello/plugins/greeter/plugin.go index d964981..916f267 100644 --- a/examples/hello/plugins/greeter/plugin.go +++ b/examples/hello/plugins/greeter/plugin.go @@ -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 } diff --git a/lagoon/commands.go b/lagoon/commands.go index 5ddb1fc..1eb9db6 100644 --- a/lagoon/commands.go +++ b/lagoon/commands.go @@ -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}) } diff --git a/lagoon/migrations.go b/lagoon/migrations.go index 2d253c4..c355e72 100644 --- a/lagoon/migrations.go +++ b/lagoon/migrations.go @@ -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,26 +91,30 @@ 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 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 { diff --git a/lagoon/migrations_test.go b/lagoon/migrations_test.go index 36466fc..e423a50 100644 --- a/lagoon/migrations_test.go +++ b/lagoon/migrations_test.go @@ -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) + } +} diff --git a/pact/capabilities.go b/pact/capabilities.go index 1567b3b..a8f44f4 100644 --- a/pact/capabilities.go +++ b/pact/capabilities.go @@ -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. diff --git a/surf/params.go b/surf/params.go new file mode 100644 index 0000000..f15e9fd --- /dev/null +++ b/surf/params.go @@ -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+":") +} diff --git a/surf/params_test.go b/surf/params_test.go new file mode 100644 index 0000000..ea5962d --- /dev/null +++ b/surf/params_test.go @@ -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") + } +} diff --git a/surf/router.go b/surf/router.go index edb0b0a..2f04741 100644 --- a/surf/router.go +++ b/surf/router.go @@ -3,7 +3,6 @@ package surf import ( "fmt" "net/http" - "strconv" "strings" "git.golem15.com/golem15/summercms/backpack" @@ -23,11 +22,12 @@ type namedMiddleware struct { } type route struct { - pluginID string - method string - path string - handler http.Handler - middleware []string + pluginID string + method string + 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) diff --git a/surf/router_test.go b/surf/router_test.go index 6a72adf..494391a 100644 --- a/surf/router_test.go +++ b/surf/router_test.go @@ -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") + } +}