package surf import ( "net/http" "regexp" "strconv" "strings" "git.golem15.com/golem15/summercms/modules/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 }