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:
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"
|
||||
@@ -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)
|
||||
|
||||
@@ -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