From a50e09ba341be7289985ab9653046ea79219e7c2 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Mon, 21 Sep 2026 19:42:20 +0200 Subject: [PATCH] fix(06-12): bound named middleware by body cap, cache factories, fail boot on bad body config and mux conflicts --- surf/bodylimit_test.go | 54 +++++++++++++++++++++++++++++++++ surf/cors_test.go | 3 ++ surf/router.go | 68 +++++++++++++++++++++++++++++++++++++----- 3 files changed, 118 insertions(+), 7 deletions(-) diff --git a/surf/bodylimit_test.go b/surf/bodylimit_test.go index 3a3fe5e..54186ee 100644 --- a/surf/bodylimit_test.go +++ b/surf/bodylimit_test.go @@ -125,3 +125,57 @@ func (p bodyEchoPlugin) Routes(r pact.Router) error { }) return nil } + +type mwReadPlugin struct{ seen *error } + +func (p mwReadPlugin) ID() string { return "golem15.mwread" } +func (p mwReadPlugin) Requires() []string { return nil } +func (p mwReadPlugin) Register(*backpack.App) error { return nil } +func (p mwReadPlugin) Boot(*backpack.App) error { return nil } +func (p mwReadPlugin) Middlewares() map[string]pact.Middleware { + return map[string]pact.Middleware{"reader": func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, *p.seen = io.ReadAll(r.Body) + next.ServeHTTP(w, r) + }) + }} +} +func (p mwReadPlugin) Routes(r pact.Router) error { + r.Group("/", Use("reader"), func(g pact.Router) { + g.Post("/x", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) + }) + return nil +} + +func TestBodyLimitBoundsNamedMiddleware(t *testing.T) { + cfg := writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n") + var seen error + h, err := Assemble(backpack.New(cfg), []party.Plugin{mwReadPlugin{seen: &seen}}) + if err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, "/x", strings.NewReader(strings.Repeat("a", 10))) + h.ServeHTTP(httptest.NewRecorder(), req) + var maxErr *http.MaxBytesError + if !errors.As(seen, &maxErr) { + t.Fatalf("middleware read err = %v, want MaxBytesError", seen) + } +} + +func TestBodyLimitMissingConfigFailsBoot(t *testing.T) { + cfg := writeHTTPConfig(t, "body_limits:\n upload_bytes: 8\n") + if _, err := BuildRouter(backpack.New(cfg), nil); err == nil || !strings.Contains(err.Error(), "default_bytes") { + t.Fatalf("err = %v", err) + } +} + +func TestCompileRouteConflictReturnsError(t *testing.T) { + r := New(nil) + ok := func(w http.ResponseWriter, _ *http.Request) {} + r.BindPlugin("a") + r.Get("/a/{x}", ok) + r.Get("/a/{y}", ok) + if _, err := r.compile(); err == nil { + t.Fatal("expected conflict error") + } +} diff --git a/surf/cors_test.go b/surf/cors_test.go index 48f2bfb..b9fdc6f 100644 --- a/surf/cors_test.go +++ b/surf/cors_test.go @@ -100,6 +100,9 @@ func writeHTTPConfig(t *testing.T, httpYAML string) *compass.Config { func TestCORSAssembleUsesConfig(t *testing.T) { cfg := writeHTTPConfig(t, ` +body_limits: + default_bytes: 1048576 + upload_bytes: 1048576 cors: paths: ["api/*"] allowed_methods: ["*"] diff --git a/surf/router.go b/surf/router.go index bf734e5..fd88af4 100644 --- a/surf/router.go +++ b/surf/router.go @@ -3,6 +3,7 @@ package surf import ( "bytes" "fmt" + "math" "net/http" "strings" "time" @@ -56,6 +57,7 @@ type Router struct { corsCfg CORSConfig defaultBytes int64 uploadBytes int64 + built map[string]pact.Middleware } var ( @@ -78,6 +80,7 @@ func New(origins []string) *Router { named: make(map[string]namedMiddleware), factories: make(map[string]namedMiddlewareFactory), seen: make(map[string]string), + built: make(map[string]pact.Middleware), origins: origins, } } @@ -345,20 +348,30 @@ func (r *Router) compile() (http.Handler, error) { if err != nil { return nil, err } - mux.Handle(rt.method+" "+rt.path, h) + if err := handleRoute(mux, rt, h); err != nil { + return nil, err + } } return pathScopedCORS(r.corsCfg, mux), nil } +// handleRoute registers a route, converting a ServeMux conflict panic into an error. +func handleRoute(mux *http.ServeMux, rt route, h http.Handler) (err error) { + defer func() { + if rec := recover(); rec != nil { + err = fmt.Errorf("surf: route conflict for %s %s (plugin %q): %v", rt.method, rt.path, rt.pluginID, rec) + } + }() + mux.Handle(rt.method+" "+rt.path, h) + return nil +} + func (r *Router) wrap(rt route) (http.Handler, error) { h := constrain(rt.handler, rt.constraints) limit, err := routeBodyLimit(rt, r.defaultBytes) if err != nil { return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err) } - if limit > 0 { - h = bodyLimit(limit)(h) - } h = orgSlot(h) for i := len(rt.middleware) - 1; i >= 0; i-- { name := rt.middleware[i] @@ -391,6 +404,10 @@ func (r *Router) wrap(rt route) (http.Handler, error) { } return nil, fmt.Errorf("surf: plugin %q references unknown middleware %q", rt.pluginID, name) } + // Body cap wraps every named/factory middleware but stays inside recovery. + if limit > 0 { + h = bodyLimit(limit)(h) + } h = locale(h) if rt.raw { h = recoverBare(h) @@ -426,7 +443,7 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) { return nil, err } if err := r.RegisterMiddlewareFactory("surf", "body.limit", func(param string) pact.Middleware { - // Limit is applied innermost in wrap(); the factory only occupies the name. + // Limit is applied outermost-inside-recovery in wrap(); the factory only occupies the name. return func(next http.Handler) http.Handler { return next } }); err != nil { return nil, err @@ -438,8 +455,15 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) { } r.corsCfg = corsCfg if app.Config != nil { - r.defaultBytes = int64(app.Config.Int("http.body_limits.default_bytes")) - r.uploadBytes = int64(app.Config.Int("http.body_limits.upload_bytes")) + d, err := requiredBytes(app, "http.body_limits.default_bytes") + if err != nil { + return nil, err + } + u, err := requiredBytes(app, "http.body_limits.upload_bytes") + if err != nil { + return nil, err + } + r.defaultBytes, r.uploadBytes = d, u } } for _, p := range plugins { @@ -494,6 +518,36 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) { return r, nil } +func requiredBytes(app *backpack.App, key string) (int64, error) { + raw, ok := app.Config.Lookup(key) + if !ok { + return 0, fmt.Errorf("surf: config %s is required", key) + } + var n int64 + switch v := raw.(type) { + case int: + n = int64(v) + case int64: + n = v + case uint64: + if v > math.MaxInt64 { + return 0, fmt.Errorf("surf: config %s out of range", key) + } + n = int64(v) + case float64: + if v != math.Trunc(v) || v > 9007199254740992 { + return 0, fmt.Errorf("surf: config %s must be a whole number", key) + } + n = int64(v) + default: + return 0, fmt.Errorf("surf: config %s must be numeric, got %T", key, raw) + } + if n < 1 { + return 0, fmt.Errorf("surf: config %s must be >= 1", key) + } + return n, nil +} + func corsOrigins(app *backpack.App) []string { if app == nil || app.Config == nil { return nil