package httpsecurity import ( "errors" "fmt" "net" "net/http" "net/netip" "strconv" "strings" ) type cidrMatcher struct { prefixes []netip.Prefix } func newCIDRMatcher(values []string) (cidrMatcher, error) { matcher := cidrMatcher{prefixes: make([]netip.Prefix, 0, len(values))} for _, value := range values { prefix, err := netip.ParsePrefix(strings.TrimSpace(value)) if err != nil { return cidrMatcher{}, fmt.Errorf("%w: invalid CIDR", ErrInvalidConfig) } matcher.prefixes = append(matcher.prefixes, prefix.Masked()) } return matcher, nil } func (matcher cidrMatcher) match(address netip.Addr) bool { address = address.Unmap() for _, prefix := range matcher.prefixes { if prefix.Contains(address) { return true } } return false } type ClientIPResolver struct { trusted cidrMatcher } func NewClientIPResolver(trustedCIDRs []string) (*ClientIPResolver, error) { trusted, err := newCIDRMatcher(trustedCIDRs) if err != nil { return nil, err } return &ClientIPResolver{trusted: trusted}, nil } func (resolver *ClientIPResolver) Resolve(request *http.Request) (netip.Addr, error) { if resolver == nil || request == nil { return netip.Addr{}, errors.New("client IP resolver and request are required") } peer, err := parseRemoteAddress(request.RemoteAddr) if err != nil { return netip.Addr{}, err } if !resolver.trusted.match(peer) { return peer, nil } chain, present, err := parseForwardedChain(request.Header.Values("Forwarded")) if err != nil { return netip.Addr{}, err } if !present { chain, err = parseXForwardedFor(request.Header.Values("X-Forwarded-For")) if err != nil { return netip.Addr{}, err } } if len(chain) == 0 { return peer, nil } for index := len(chain) - 1; index >= 0; index-- { if !resolver.trusted.match(chain[index]) { return chain[index], nil } } return chain[0], nil } func parseRemoteAddress(remote string) (netip.Addr, error) { host, _, err := net.SplitHostPort(strings.TrimSpace(remote)) if err != nil { host = strings.TrimSpace(remote) } address, err := netip.ParseAddr(host) if err != nil { return netip.Addr{}, errors.New("invalid remote address") } return address.Unmap(), nil } func parseXForwardedFor(fields []string) ([]netip.Addr, error) { chain := make([]netip.Addr, 0, len(fields)+1) for _, field := range fields { for value := range strings.SplitSeq(field, ",") { address, err := netip.ParseAddr(strings.TrimSpace(value)) if err != nil { return nil, errors.New("invalid X-Forwarded-For address") } chain = append(chain, address.Unmap()) } } return chain, nil } func parseForwardedChain(fields []string) ([]netip.Addr, bool, error) { if len(fields) == 0 { return nil, false, nil } chain := make([]netip.Addr, 0, len(fields)+1) for _, field := range fields { for element := range strings.SplitSeq(field, ",") { found := false for parameter := range strings.SplitSeq(element, ";") { name, value, ok := strings.Cut(strings.TrimSpace(parameter), "=") if !ok || !strings.EqualFold(name, "for") { continue } address, err := parseForwardedIdentifier(value) if err != nil { return nil, true, err } chain = append(chain, address) found = true break } if !found { return nil, true, errors.New("Forwarded element is missing for parameter") } } } return chain, true, nil } func parseForwardedIdentifier(raw string) (netip.Addr, error) { value := strings.TrimSpace(raw) if strings.HasPrefix(value, `"`) { unquoted, err := strconv.Unquote(value) if err != nil { return netip.Addr{}, errors.New("invalid quoted Forwarded identifier") } value = unquoted } if strings.EqualFold(value, "unknown") || strings.HasPrefix(value, "_") { return netip.Addr{}, errors.New("non-IP Forwarded identifier") } if strings.HasPrefix(value, "[") { if addressPort, err := netip.ParseAddrPort(value); err == nil { return addressPort.Addr().Unmap(), nil } if !strings.HasSuffix(value, "]") { return netip.Addr{}, errors.New("invalid Forwarded IPv6 identifier") } value = strings.TrimSuffix(strings.TrimPrefix(value, "["), "]") } else if host, _, err := net.SplitHostPort(value); err == nil { value = host } address, err := netip.ParseAddr(value) if err != nil { return netip.Addr{}, errors.New("invalid Forwarded address") } return address.Unmap(), nil }