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
This commit is contained in:
46
surf/bodylimit.go
Normal file
46
surf/bodylimit.go
Normal file
@@ -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
|
||||||
|
}
|
||||||
127
surf/bodylimit_test.go
Normal file
127
surf/bodylimit_test.go
Normal file
@@ -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
|
||||||
|
}
|
||||||
161
surf/cors.go
Normal file
161
surf/cors.go
Normal file
@@ -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
|
||||||
|
}
|
||||||
119
surf/cors_test.go
Normal file
119
surf/cors_test.go
Normal file
@@ -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"))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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 {
|
if err := r.RegisterMiddleware("golem15.user", "jwt.auth", record("jwt.auth")); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -51,6 +51,9 @@ type Router struct {
|
|||||||
origins []string
|
origins []string
|
||||||
compileErr error
|
compileErr error
|
||||||
limiter *FixedWindowLimiter
|
limiter *FixedWindowLimiter
|
||||||
|
corsCfg CORSConfig
|
||||||
|
defaultBytes int64
|
||||||
|
uploadBytes int64
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -342,11 +345,18 @@ func (r *Router) compile() (http.Handler, error) {
|
|||||||
}
|
}
|
||||||
mux.Handle(rt.method+" "+rt.path, h)
|
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) {
|
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)
|
||||||
|
if err != nil {
|
||||||
|
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]
|
||||||
@@ -368,6 +378,11 @@ func (r *Router) wrap(rt route) (http.Handler, error) {
|
|||||||
return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err)
|
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)
|
h = factory.fn(param)(h)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -408,6 +423,23 @@ func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) {
|
|||||||
}); err != nil {
|
}); err != nil {
|
||||||
return nil, err
|
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 {
|
for _, p := range plugins {
|
||||||
if hm, ok := p.(pact.HasMiddleware); ok {
|
if hm, ok := p.(pact.HasMiddleware); ok {
|
||||||
for name, fn := range hm.Middlewares() {
|
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 {
|
func locale(next http.Handler) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
next.ServeHTTP(w, r.WithContext(towel.WithLocale(r.Context(), r.Header.Get("Accept-Language"))))
|
next.ServeHTTP(w, r.WithContext(towel.WithLocale(r.Context(), r.Header.Get("Accept-Language"))))
|
||||||
|
|||||||
@@ -143,7 +143,13 @@ func TestRecoverReturnsOpaqueJSON500(t *testing.T) {
|
|||||||
|
|
||||||
func TestCORSPreflightBypassesNamedAuth(t *testing.T) {
|
func TestCORSPreflightBypassesNamedAuth(t *testing.T) {
|
||||||
called := false
|
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 {
|
if err := r.RegisterMiddleware("golem15.user", "jwt.auth", func(next http.Handler) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||||
called = true
|
called = true
|
||||||
|
|||||||
Reference in New Issue
Block a user