Files
summercms/surf/cors.go
Jakub Zych 30539b954f 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
2026-09-19 20:10:02 +02:00

162 lines
4.2 KiB
Go

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
}