From 30539b954fa1f4db82b269c233389bfae60dec21 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Sat, 19 Sep 2026 20:10:02 +0200 Subject: [PATCH] feat(06-03): add path-scoped CORS and per-route body limits - CORS matches Laravel path globs (api/* includes nested segments); unlisted paths get no headers - Non-raw routes wrap http.MaxBytesReader from http.body_limits.default_bytes; body.limit:N overrides innermost - Raw routes stay uncapped at this layer --- surf/bodylimit.go | 46 ++++++++++++ surf/bodylimit_test.go | 127 +++++++++++++++++++++++++++++++ surf/cors.go | 161 ++++++++++++++++++++++++++++++++++++++++ surf/cors_test.go | 119 +++++++++++++++++++++++++++++ surf/middleware_test.go | 8 +- surf/router.go | 75 +++++++++++-------- surf/router_test.go | 8 +- 7 files changed, 510 insertions(+), 34 deletions(-) create mode 100644 surf/bodylimit.go create mode 100644 surf/bodylimit_test.go create mode 100644 surf/cors.go create mode 100644 surf/cors_test.go diff --git a/surf/bodylimit.go b/surf/bodylimit.go new file mode 100644 index 0000000..dc533e8 --- /dev/null +++ b/surf/bodylimit.go @@ -0,0 +1,46 @@ +package surf + +import ( + "fmt" + "net/http" + "strconv" + "strings" +) + +func bodyLimit(n int64) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if n > 0 && r.Body != nil { + r.Body = http.MaxBytesReader(w, r.Body, n) + } + next.ServeHTTP(w, r) + }) + } +} + +func parseBodyLimit(param string) (int64, error) { + n, err := strconv.ParseInt(param, 10, 64) + if err != nil || n <= 0 { + return 0, fmt.Errorf("invalid body.limit %q", param) + } + return n, nil +} + +func routeBodyLimit(rt route, defaultBytes int64) (int64, error) { + if rt.raw { + return 0, nil + } + limit := defaultBytes + for _, name := range rt.middleware { + base, param, ok := strings.Cut(name, ":") + if !ok || base != "body.limit" { + continue + } + n, err := parseBodyLimit(param) + if err != nil { + return 0, err + } + limit = n + } + return limit, nil +} diff --git a/surf/bodylimit_test.go b/surf/bodylimit_test.go new file mode 100644 index 0000000..3a3fe5e --- /dev/null +++ b/surf/bodylimit_test.go @@ -0,0 +1,127 @@ +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 +} diff --git a/surf/cors.go b/surf/cors.go new file mode 100644 index 0000000..a9ab1ee --- /dev/null +++ b/surf/cors.go @@ -0,0 +1,161 @@ +package surf + +import ( + "net/http" + "regexp" + "strconv" + "strings" + + "git.golem15.com/golem15/summercms/compass" +) + +// CORSConfig matches Laravel config/cors.php keys. +type CORSConfig struct { + Paths []string `koanf:"paths"` + AllowedMethods []string `koanf:"allowed_methods"` + AllowedOrigins []string `koanf:"allowed_origins"` + AllowedOriginsPatterns []string `koanf:"allowed_origins_patterns"` + AllowedHeaders []string `koanf:"allowed_headers"` + ExposedHeaders []string `koanf:"exposed_headers"` + MaxAge int `koanf:"max_age"` + SupportsCredentials bool `koanf:"supports_credentials"` +} + +// LoadCORSConfig reads http.cors. Missing section yields a zero config (no +// path matches, so no CORS headers). +func LoadCORSConfig(cfg *compass.Config) (CORSConfig, error) { + var out CORSConfig + if cfg == nil || !cfg.Has("http.cors") { + return out, nil + } + if err := cfg.LoadSection("http.cors", &out); err != nil { + return out, err + } + return out, nil +} + +func pathScopedCORS(cfg CORSConfig, next http.Handler) http.Handler { + globs := make([]*regexp.Regexp, 0, len(cfg.Paths)) + for _, p := range cfg.Paths { + if re := compileLaravelGlob(p); re != nil { + globs = append(globs, re) + } + } + originPats := make([]*regexp.Regexp, 0, len(cfg.AllowedOriginsPatterns)) + for _, p := range cfg.AllowedOriginsPatterns { + re, err := regexp.Compile(p) + if err != nil { + continue + } + originPats = append(originPats, re) + } + allowAnyOrigin := containsStar(cfg.AllowedOrigins) + allowAnyMethod := containsStar(cfg.AllowedMethods) + allowAnyHeader := containsStar(cfg.AllowedHeaders) + origins := make(map[string]struct{}, len(cfg.AllowedOrigins)) + for _, o := range cfg.AllowedOrigins { + if o != "" && o != "*" { + origins[o] = struct{}{} + } + } + methods := strings.Join(cfg.AllowedMethods, ", ") + headers := strings.Join(cfg.AllowedHeaders, ", ") + exposed := strings.Join(cfg.ExposedHeaders, ", ") + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !pathMatchesCORS(globs, r.URL.Path) { + next.ServeHTTP(w, r) + return + } + origin := r.Header.Get("Origin") + allowed, value := corsAllowOrigin(allowAnyOrigin, origins, originPats, origin) + if allowed { + w.Header().Set("Access-Control-Allow-Origin", value) + if value != "*" { + w.Header().Set("Vary", "Origin") + } + if allowAnyMethod { + w.Header().Set("Access-Control-Allow-Methods", "*") + } else if methods != "" { + w.Header().Set("Access-Control-Allow-Methods", methods) + } + if allowAnyHeader { + w.Header().Set("Access-Control-Allow-Headers", "*") + } else if headers != "" { + w.Header().Set("Access-Control-Allow-Headers", headers) + } + if exposed != "" { + w.Header().Set("Access-Control-Expose-Headers", exposed) + } + if cfg.MaxAge > 0 { + w.Header().Set("Access-Control-Max-Age", strconv.Itoa(cfg.MaxAge)) + } + if cfg.SupportsCredentials { + w.Header().Set("Access-Control-Allow-Credentials", "true") + } + } + if r.Method == http.MethodOptions { + w.WriteHeader(http.StatusNoContent) + return + } + next.ServeHTTP(w, r) + }) +} + +func corsAllowOrigin(allowAny bool, origins map[string]struct{}, pats []*regexp.Regexp, origin string) (bool, string) { + if allowAny { + return true, "*" + } + if origin == "" { + return false, "" + } + if _, ok := origins[origin]; ok { + return true, origin + } + for _, re := range pats { + if re.MatchString(origin) { + return true, origin + } + } + return false, "" +} + +func pathMatchesCORS(globs []*regexp.Regexp, urlPath string) bool { + trimmed := strings.TrimPrefix(urlPath, "/") + for _, re := range globs { + if re.MatchString(trimmed) { + return true + } + } + return false +} + +func compileLaravelGlob(pattern string) *regexp.Regexp { + pattern = strings.TrimPrefix(pattern, "/") + var b strings.Builder + b.WriteString("(?s)^") + for i := 0; i < len(pattern); i++ { + switch pattern[i] { + case '*': + b.WriteString(".*") + case '?': + b.WriteByte('.') + default: + b.WriteString(regexp.QuoteMeta(pattern[i : i+1])) + } + } + b.WriteByte('$') + re, err := regexp.Compile(b.String()) + if err != nil { + return nil + } + return re +} + +func containsStar(vals []string) bool { + for _, v := range vals { + if v == "*" { + return true + } + } + return false +} diff --git a/surf/cors_test.go b/surf/cors_test.go new file mode 100644 index 0000000..48f2bfb --- /dev/null +++ b/surf/cors_test.go @@ -0,0 +1,119 @@ +package surf + +import ( + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "testing" + + "git.golem15.com/golem15/summercms/backpack" + "git.golem15.com/golem15/summercms/compass" + "git.golem15.com/golem15/summercms/party" +) + +func TestCORSLaravelGlobMatchesNestedPaths(t *testing.T) { + re := compileLaravelGlob("api/*") + if re == nil || !re.MatchString("api/v1/fonoteka/genres") { + t.Fatal("api/* must match api/v1/fonoteka/genres (Laravel Str::is, not Go path.Match)") + } + if re.MatchString("_fonoteka/api/v1/genres") { + t.Fatal("api/* must not match _fonoteka/api/v1/genres") + } + mcp := compileLaravelGlob("oauth/mcp/*") + if mcp == nil || !mcp.MatchString("oauth/mcp/token") { + t.Fatal("oauth/mcp/* must match oauth/mcp/token") + } +} + +func TestCORSPathScopedHeaders(t *testing.T) { + cfg := CORSConfig{ + Paths: []string{"api/*", "oauth/mcp/*"}, + AllowedMethods: []string{"*"}, + AllowedOrigins: []string{"*"}, + AllowedHeaders: []string{"*"}, + } + mux := http.NewServeMux() + mux.HandleFunc("GET /api/v1/fonoteka/genres", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + mux.HandleFunc("GET /_fonoteka/api/v1/genres", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + h := pathScopedCORS(cfg, mux) + + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/v1/fonoteka/genres", nil)) + if rec.Header().Get("Access-Control-Allow-Origin") != "*" { + t.Fatalf("token group ACAO = %q", rec.Header().Get("Access-Control-Allow-Origin")) + } + + rec = httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/_fonoteka/api/v1/genres", nil)) + if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" { + t.Fatalf("JWT group ACAO = %q, want empty", got) + } +} + +func TestCORSConfigLoadedFromHTTPSection(t *testing.T) { + cfg := writeHTTPConfig(t, ` +cors: + paths: ["api/*"] + allowed_methods: ["*"] + allowed_origins: ["*"] + allowed_origins_patterns: [] + allowed_headers: ["*"] + exposed_headers: [] + max_age: 0 + supports_credentials: false +`) + got, err := LoadCORSConfig(cfg) + if err != nil { + t.Fatal(err) + } + if len(got.Paths) != 1 || got.Paths[0] != "api/*" { + t.Fatalf("paths = %v", got.Paths) + } + if !containsStar(got.AllowedOrigins) { + t.Fatalf("origins = %v", got.AllowedOrigins) + } +} + +func writeHTTPConfig(t *testing.T, httpYAML string) *compass.Config { + t.Helper() + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: cors-test\n"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte(httpYAML), 0o644); err != nil { + t.Fatal(err) + } + cfg, err := compass.Open(compass.Options{ + Dir: dir, + Environ: []string{"SUMMER_ENV=development"}, + }) + if err != nil { + t.Fatal(err) + } + return cfg +} + +func TestCORSAssembleUsesConfig(t *testing.T) { + cfg := writeHTTPConfig(t, ` +cors: + paths: ["api/*"] + allowed_methods: ["*"] + allowed_origins: ["*"] + allowed_headers: ["*"] +`) + p := assemblePlugin{id: "golem15.demo", path: "/genres"} + h, err := Assemble(backpack.New(cfg), []party.Plugin{p}) + if err != nil { + t.Fatal(err) + } + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/genres", nil)) + if rec.Header().Get("Access-Control-Allow-Origin") != "*" { + t.Fatalf("ACAO = %q", rec.Header().Get("Access-Control-Allow-Origin")) + } +} diff --git a/surf/middleware_test.go b/surf/middleware_test.go index 0c807b8..84b6543 100644 --- a/surf/middleware_test.go +++ b/surf/middleware_test.go @@ -91,7 +91,13 @@ func TestPipelineOrderRecoverCORSLocaleAuthPasswordOrgRateHandler(t *testing.T) } } - r := New([]string{"http://localhost:3000"}) + r := New(nil) + r.corsCfg = CORSConfig{ + Paths: []string{"api/*"}, + AllowedOrigins: []string{"http://localhost:3000"}, + AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"}, + AllowedHeaders: []string{"Authorization", "Content-Type", "Accept"}, + } if err := r.RegisterMiddleware("golem15.user", "jwt.auth", record("jwt.auth")); err != nil { t.Fatal(err) } diff --git a/surf/router.go b/surf/router.go index b274aa7..a35ade5 100644 --- a/surf/router.go +++ b/surf/router.go @@ -41,16 +41,19 @@ type route struct { // Router compiles group declarations onto net/http ServeMux. type Router struct { - pluginID string - prefix string - middleware []string - named map[string]namedMiddleware - factories map[string]namedMiddlewareFactory - routes []route - seen map[string]string - origins []string - compileErr error - limiter *FixedWindowLimiter + pluginID string + prefix string + middleware []string + named map[string]namedMiddleware + factories map[string]namedMiddlewareFactory + routes []route + seen map[string]string + origins []string + compileErr error + limiter *FixedWindowLimiter + corsCfg CORSConfig + defaultBytes int64 + uploadBytes int64 } var ( @@ -342,11 +345,18 @@ func (r *Router) compile() (http.Handler, error) { } mux.Handle(rt.method+" "+rt.path, h) } - return cors(r.origins, mux), nil + return pathScopedCORS(r.corsCfg, mux), 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] @@ -368,6 +378,11 @@ func (r *Router) wrap(rt route) (http.Handler, error) { return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err) } } + if base == "body.limit" { + if _, err := parseBodyLimit(param); err != nil { + return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err) + } + } h = factory.fn(param)(h) continue } @@ -408,6 +423,23 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) { }); err != nil { 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. + return func(next http.Handler) http.Handler { return next } + }); err != nil { + return nil, err + } + if app != nil { + corsCfg, err := LoadCORSConfig(app.Config) + if err != nil { + return nil, err + } + 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")) + } + } for _, p := range plugins { if hm, ok := p.(pact.HasMiddleware); ok { for name, fn := range hm.Middlewares() { @@ -509,27 +541,6 @@ func recoverBare(next http.Handler) http.Handler { }) } -func cors(origins []string, next http.Handler) http.Handler { - allowed := make(map[string]struct{}, len(origins)) - for _, o := range origins { - allowed[o] = struct{}{} - } - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - origin := r.Header.Get("Origin") - if _, ok := allowed[origin]; ok && origin != "" { - w.Header().Set("Access-Control-Allow-Origin", origin) - w.Header().Set("Vary", "Origin") - w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, Accept") - w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS") - } - if r.Method == http.MethodOptions { - w.WriteHeader(http.StatusNoContent) - return - } - next.ServeHTTP(w, r) - }) -} - func locale(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { next.ServeHTTP(w, r.WithContext(towel.WithLocale(r.Context(), r.Header.Get("Accept-Language")))) diff --git a/surf/router_test.go b/surf/router_test.go index daec18b..fa5b118 100644 --- a/surf/router_test.go +++ b/surf/router_test.go @@ -143,7 +143,13 @@ func TestRecoverReturnsOpaqueJSON500(t *testing.T) { func TestCORSPreflightBypassesNamedAuth(t *testing.T) { called := false - r := New([]string{"http://localhost:3000"}) + r := New(nil) + r.corsCfg = CORSConfig{ + Paths: []string{"api/*"}, + AllowedOrigins: []string{"http://localhost:3000"}, + AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"}, + AllowedHeaders: []string{"Authorization", "Content-Type", "Accept"}, + } if err := r.RegisterMiddleware("golem15.user", "jwt.auth", func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { called = true