package surf import ( "errors" "io" "net/http" "net/http/httptest" "strings" "testing" "git.golem15.com/golem15/summercms/backpack" "git.golem15.com/golem15/summercms/pact" "git.golem15.com/golem15/summercms/party" ) func TestBodyLimitDefaultRejectsOversizedBody(t *testing.T) { cfg := writeHTTPConfig(t, ` body_limits: default_bytes: 32 upload_bytes: 64 `) p := bodyEchoPlugin{id: "golem15.demo"} h, err := Assemble(backpack.New(cfg), []party.Plugin{p}) if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader(strings.Repeat("a", 64))) h.ServeHTTP(rec, req) if rec.Code != http.StatusRequestEntityTooLarge { t.Fatalf("status = %d body=%q (want 413 from MaxBytesReader)", rec.Code, rec.Body.String()) } rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader("ok")) h.ServeHTTP(rec, req) if rec.Code != http.StatusOK { t.Fatalf("small body status = %d", rec.Code) } } func TestBodyLimitRawExempt(t *testing.T) { cfg := writeHTTPConfig(t, ` body_limits: default_bytes: 8 upload_bytes: 8 `) p := bodyEchoPlugin{id: "golem15.demo", raw: true} h, err := Assemble(backpack.New(cfg), []party.Plugin{p}) if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader(strings.Repeat("a", 64))) h.ServeHTTP(rec, req) if rec.Code != http.StatusOK { t.Fatalf("raw route should not apply default body limit, status = %d", rec.Code) } } func TestBodyLimitOverride(t *testing.T) { cfg := writeHTTPConfig(t, ` body_limits: default_bytes: 8 upload_bytes: 64 `) p := bodyEchoPlugin{id: "golem15.demo", extra: []string{"body.limit:64"}} h, err := Assemble(backpack.New(cfg), []party.Plugin{p}) if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader(strings.Repeat("a", 32))) h.ServeHTTP(rec, req) if rec.Code != http.StatusOK { t.Fatalf("override should allow 32 bytes, status = %d", rec.Code) } } func TestBodyLimitLoadsConfigValues(t *testing.T) { cfg := writeHTTPConfig(t, ` body_limits: default_bytes: 8388608 upload_bytes: 2097152 `) r, err := BuildRouter(backpack.New(cfg), nil) if err != nil { t.Fatal(err) } if r.defaultBytes != 8388608 || r.uploadBytes != 2097152 { t.Fatalf("limits = %d / %d", r.defaultBytes, r.uploadBytes) } } type bodyEchoPlugin struct { id string raw bool extra []string } func (p bodyEchoPlugin) ID() string { return p.id } func (p bodyEchoPlugin) Requires() []string { return nil } func (p bodyEchoPlugin) Register(*backpack.App) error { return nil } func (p bodyEchoPlugin) Boot(*backpack.App) error { return nil } func (p bodyEchoPlugin) Routes(r pact.Router) error { h := func(w http.ResponseWriter, req *http.Request) { _, err := io.Copy(io.Discard, req.Body) if err != nil { var maxErr *http.MaxBytesError if errors.As(err, &maxErr) { w.WriteHeader(http.StatusRequestEntityTooLarge) return } w.WriteHeader(http.StatusInternalServerError) return } w.WriteHeader(http.StatusOK) } open := r.Group if p.raw { open = r.GroupRaw } open("/", Use(p.extra...), func(g pact.Router) { g.Post("/echo", h) }) 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") } } type bodyProbePlugin struct { id string raw bool use []string mw map[string]pact.Middleware routes func(pact.Router) } func (p bodyProbePlugin) ID() string { return p.id } func (p bodyProbePlugin) Requires() []string { return nil } func (p bodyProbePlugin) Register(*backpack.App) error { return nil } func (p bodyProbePlugin) Boot(*backpack.App) error { return nil } func (p bodyProbePlugin) Middlewares() map[string]pact.Middleware { return p.mw } func (p bodyProbePlugin) Routes(r pact.Router) error { open := r.Group if p.raw { open = r.GroupRaw } open("/", Use(p.use...), func(g pact.Router) { g.Post("/x", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) }) return nil } func TestBodyLimitBoundsBodyConsumingMiddleware(t *testing.T) { type result struct { n int err error } run := func(t *testing.T, raw bool, use []string, bodyLen int, after func()) (result, *httptest.ResponseRecorder) { t.Helper() var res result reader := func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { b, err := io.ReadAll(r.Body) res = result{n: len(b), err: err} if after != nil { after() } next.ServeHTTP(w, r) }) } cfg := writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n") h, err := Assemble(backpack.New(cfg), []party.Plugin{bodyProbePlugin{ id: "golem15.probe", raw: raw, use: use, mw: map[string]pact.Middleware{"reader": reader}, }}) if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/x", strings.NewReader(strings.Repeat("a", bodyLen)))) return res, rec } var maxErr *http.MaxBytesError t.Run("default limit", func(t *testing.T) { res, _ := run(t, false, []string{"reader"}, 10, nil) if !errors.As(res.err, &maxErr) || res.n > 4 { t.Fatalf("read %d bytes, err = %v", res.n, res.err) } }) t.Run("body.limit override raises cap", func(t *testing.T) { res, _ := run(t, false, []string{"reader", "body.limit:6"}, 5, nil) if res.err != nil || res.n != 5 { t.Fatalf("read %d bytes, err = %v", res.n, res.err) } }) t.Run("body.limit override still bounds", func(t *testing.T) { res, _ := run(t, false, []string{"reader", "body.limit:6"}, 10, nil) if !errors.As(res.err, &maxErr) || res.n > 6 { t.Fatalf("read %d bytes, err = %v", res.n, res.err) } }) t.Run("raw route unaffected", func(t *testing.T) { res, rec := run(t, true, []string{"reader"}, 10, nil) if res.err != nil || res.n != 10 || rec.Code != http.StatusOK { t.Fatalf("read %d bytes, err = %v, code %d", res.n, res.err, rec.Code) } }) t.Run("panic after read still clean 500", func(t *testing.T) { _, rec := run(t, false, []string{"reader"}, 10, func() { panic("boom") }) if rec.Code != http.StatusInternalServerError || strings.Contains(rec.Body.String(), "boom") { t.Fatalf("code %d body %q", rec.Code, rec.Body.String()) } }) } func TestBodyLimitInvalidParamFailsBoot(t *testing.T) { for _, param := range []string{"abc", "0", "-5", ""} { t.Run(param, func(t *testing.T) { cfg := writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n") _, err := Assemble(backpack.New(cfg), []party.Plugin{bodyProbePlugin{ id: "golem15.probe", use: []string{"body.limit:" + param}, }}) if err == nil || !strings.Contains(err.Error(), "body.limit") { t.Fatalf("err = %v", err) } }) } }