168 lines
4.3 KiB
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
|
|
}
|