fix(06-12): bound named middleware by body cap, cache factories, fail boot on bad body config and mux conflicts
This commit is contained in:
@@ -125,3 +125,57 @@ func (p bodyEchoPlugin) Routes(r pact.Router) error {
|
|||||||
})
|
})
|
||||||
return nil
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -100,6 +100,9 @@ func writeHTTPConfig(t *testing.T, httpYAML string) *compass.Config {
|
|||||||
|
|
||||||
func TestCORSAssembleUsesConfig(t *testing.T) {
|
func TestCORSAssembleUsesConfig(t *testing.T) {
|
||||||
cfg := writeHTTPConfig(t, `
|
cfg := writeHTTPConfig(t, `
|
||||||
|
body_limits:
|
||||||
|
default_bytes: 1048576
|
||||||
|
upload_bytes: 1048576
|
||||||
cors:
|
cors:
|
||||||
paths: ["api/*"]
|
paths: ["api/*"]
|
||||||
allowed_methods: ["*"]
|
allowed_methods: ["*"]
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package surf
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"math"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -56,6 +57,7 @@ type Router struct {
|
|||||||
corsCfg CORSConfig
|
corsCfg CORSConfig
|
||||||
defaultBytes int64
|
defaultBytes int64
|
||||||
uploadBytes int64
|
uploadBytes int64
|
||||||
|
built map[string]pact.Middleware
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -78,6 +80,7 @@ func New(origins []string) *Router {
|
|||||||
named: make(map[string]namedMiddleware),
|
named: make(map[string]namedMiddleware),
|
||||||
factories: make(map[string]namedMiddlewareFactory),
|
factories: make(map[string]namedMiddlewareFactory),
|
||||||
seen: make(map[string]string),
|
seen: make(map[string]string),
|
||||||
|
built: make(map[string]pact.Middleware),
|
||||||
origins: origins,
|
origins: origins,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -345,20 +348,30 @@ func (r *Router) compile() (http.Handler, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
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
|
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) {
|
func (r *Router) wrap(rt route) (http.Handler, error) {
|
||||||
h := constrain(rt.handler, rt.constraints)
|
h := constrain(rt.handler, rt.constraints)
|
||||||
limit, err := routeBodyLimit(rt, r.defaultBytes)
|
limit, err := routeBodyLimit(rt, r.defaultBytes)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err)
|
return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err)
|
||||||
}
|
}
|
||||||
if limit > 0 {
|
|
||||||
h = bodyLimit(limit)(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-- {
|
||||||
name := rt.middleware[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)
|
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)
|
h = locale(h)
|
||||||
if rt.raw {
|
if rt.raw {
|
||||||
h = recoverBare(h)
|
h = recoverBare(h)
|
||||||
@@ -426,7 +443,7 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := r.RegisterMiddlewareFactory("surf", "body.limit", func(param string) pact.Middleware {
|
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 }
|
return func(next http.Handler) http.Handler { return next }
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -438,8 +455,15 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) {
|
|||||||
}
|
}
|
||||||
r.corsCfg = corsCfg
|
r.corsCfg = corsCfg
|
||||||
if app.Config != nil {
|
if app.Config != nil {
|
||||||
r.defaultBytes = int64(app.Config.Int("http.body_limits.default_bytes"))
|
d, err := requiredBytes(app, "http.body_limits.default_bytes")
|
||||||
r.uploadBytes = int64(app.Config.Int("http.body_limits.upload_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 {
|
for _, p := range plugins {
|
||||||
@@ -494,6 +518,36 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) {
|
|||||||
return r, nil
|
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 {
|
func corsOrigins(app *backpack.App) []string {
|
||||||
if app == nil || app.Config == nil {
|
if app == nil || app.Config == nil {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
Reference in New Issue
Block a user