// Copyright (c) 2025-2026 Probo Inc . // // Permission is hereby granted, free of charge, to any person obtaining a copy // of this software and associated documentation files (the "Software"), to deal // in the Software without restriction, including without limitation the rights // to use, copy, modify, merge, publish, distribute, sublicense, and/or sell // copies of the Software, and to permit persons to whom the Software is // furnished to do so, subject to the following conditions: // // The above copyright notice and this permission notice shall be included in // all copies or substantial portions of the Software. // // THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR // IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, // FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE // AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER // LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, // OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE // SOFTWARE. package saferedirect import ( "context" "net/http" "net/url" "path" "strings" ) type ( AllowedHostFunc func(ctx context.Context, host string) bool SafeRedirect struct { allowedHost AllowedHostFunc } ) func New(allowedHost AllowedHostFunc) *SafeRedirect { return &SafeRedirect{allowedHost: allowedHost} } // StaticHosts returns an AllowedHost function that matches against a fixed // list of hosts. func StaticHosts(hosts ...string) AllowedHostFunc { allowed := make(map[string]bool, len(hosts)) for _, h := range hosts { if h != "" { allowed[h] = true } } return func(_ context.Context, host string) bool { return allowed[host] } } // Origins returns an AllowedHost function that matches hosts extracted from // absolute origin URLs (for example CORS allowed-origins). Invalid or empty // origins are ignored. The host includes the port when present. func Origins(origins ...string) AllowedHostFunc { hosts := make([]string, 0, len(origins)) for _, origin := range origins { parsed, err := url.Parse(origin) if err != nil || parsed.Host == "" { continue } hosts = append(hosts, parsed.Host) } return StaticHosts(hosts...) } // Any returns an AllowedHost function that allows a host when any of the // provided functions allow it. Nil functions are skipped. func Any(fns ...AllowedHostFunc) AllowedHostFunc { return func(ctx context.Context, host string) bool { for _, fn := range fns { if fn != nil && fn(ctx, host) { return true } } return false } } func (sr *SafeRedirect) Validate(ctx context.Context, redirectURL string) (string, bool) { if redirectURL == "" { return "", false } if strings.HasPrefix(redirectURL, "/") { safePath, ok := normalizeRelativePath(redirectURL) if !ok { return "", false } return safePath, true } parsedURL, err := url.Parse(redirectURL) if err != nil { return "", false } if parsedURL.IsAbs() { if parsedURL.Scheme != "http" && parsedURL.Scheme != "https" { return "", false } if sr.allowedHost != nil && !sr.allowedHost(ctx, parsedURL.Host) { return "", false } return redirectURL, true } return "", false } func (sr *SafeRedirect) GetSafeRedirectURL(ctx context.Context, redirectURL, fallbackURL string) string { if safeURL, isValid := sr.Validate(ctx, redirectURL); isValid { return safeURL } return fallbackURL } func (sr *SafeRedirect) Redirect(w http.ResponseWriter, r *http.Request, redirectURL, fallbackURL string, statusCode int) { safeURL := sr.GetSafeRedirectURL(r.Context(), redirectURL, fallbackURL) http.Redirect(w, r, safeURL, statusCode) } func normalizeRelativePath(redirectURL string) (string, bool) { if strings.HasPrefix(redirectURL, "//") { return "", false } if strings.Contains(redirectURL, `\`) || strings.Contains(strings.ToLower(redirectURL), "%5c") { return "", false } cleaned := path.Clean(redirectURL) if !strings.HasPrefix(cleaned, "/") { return "", false } if len(cleaned) > 1 && (cleaned[1] == '/' || cleaned[1] == '\\') { return "", false } if strings.Contains(cleaned, `\`) { return "", false } return cleaned, true }