proxy-pool/internal/platform/httpsecurity/source.go

168 lines
4.3 KiB
Go

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
}