package main import ( "log" "net/http" "strings" "golang.org/x/time/rate" ) // RateLimiter middleware limits requests per IP func RateLimiter(rps float64, burst int) func(http.Handler) http.Handler { limiter := rate.NewLimiter(rate.Limit(rps), burst) return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { if !limiter.Allow() { http.Error(w, http.StatusText(http.StatusTooManyRequests), http.StatusTooManyRequests) return } next.ServeHTTP(w, req) }) } } func RefererCheck(allowedOrigins []string) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { referer := r.Header.Get("Referer") origin := r.Header.Get("Origin") // Debug logging log.Printf("RefererCheck - Referer: '%s', Origin: '%s', AllowedOrigins: %+v", referer, origin, allowedOrigins) log.Printf("Request headers: %+v", r.Header) // Check both Referer and Origin headers if referer != "" { for _, allowedOrigin := range allowedOrigins { if strings.HasPrefix(referer, strings.TrimSuffix(allowedOrigin, "/")) { next.ServeHTTP(w, r) return } } } if origin != "" { for _, allowedOrigin := range allowedOrigins { if origin == strings.TrimSuffix(allowedOrigin, "/") { next.ServeHTTP(w, r) return } } } log.Printf("RefererCheck - BLOCKED: No valid referer or origin found") http.Error(w, "Forbidden", http.StatusForbidden) }) } }