feat(ssrf): implement SSRF protection in HTTP clients and validation functions

This commit is contained in:
CaIon
2026-07-06 14:52:01 +08:00
parent 1e11dfcfb5
commit df087b022d
10 changed files with 799 additions and 97 deletions
+97 -61
View File
@@ -29,24 +29,42 @@ var DefaultSSRFProtection = &SSRFProtection{
AllowedPorts: []int{},
}
// NewSSRFProtectionFromFetchSetting builds an SSRFProtection from persisted fetch_setting values.
func NewSSRFProtectionFromFetchSetting(allowPrivateIp bool, domainFilterMode bool, ipFilterMode bool, domainList, ipList, allowedPorts []string, applyIPFilterForDomain bool) (*SSRFProtection, error) {
allowedPortInts, err := parsePortRanges(allowedPorts)
if err != nil {
return nil, fmt.Errorf("request reject - invalid port configuration: %v", err)
}
return &SSRFProtection{
AllowPrivateIp: allowPrivateIp,
DomainFilterMode: domainFilterMode,
DomainList: domainList,
IpFilterMode: ipFilterMode,
IpList: ipList,
AllowedPorts: allowedPortInts,
ApplyIPFilterForDomain: applyIPFilterForDomain,
}, nil
}
// privateIPv4Nets IPv4 私有/保留/特殊用途网段
// 参考 IANA IPv4 Special-Purpose Address Registry
// https://www.iana.org/assignments/iana-ipv4-special-registry/
var privateIPv4Nets = []net.IPNet{
{IP: net.IPv4(0, 0, 0, 0), Mask: net.CIDRMask(8, 32)}, // 0.0.0.0/8 ("This network" / 未指定)
{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}, // 10.0.0.0/8 (私有)
{IP: net.IPv4(100, 64, 0, 0), Mask: net.CIDRMask(10, 32)}, // 100.64.0.0/10 (运营商级 NAT / CGNAT)
{IP: net.IPv4(127, 0, 0, 0), Mask: net.CIDRMask(8, 32)}, // 127.0.0.0/8 (回环)
{IP: net.IPv4(169, 254, 0, 0), Mask: net.CIDRMask(16, 32)}, // 169.254.0.0/16 (链路本地)
{IP: net.IPv4(172, 16, 0, 0), Mask: net.CIDRMask(12, 32)}, // 172.16.0.0/12 (私有)
{IP: net.IPv4(192, 0, 0, 0), Mask: net.CIDRMask(24, 32)}, // 192.0.0.0/24 (IETF 协议分配)
{IP: net.IPv4(192, 0, 2, 0), Mask: net.CIDRMask(24, 32)}, // 192.0.2.0/24 (TEST-NET-1)
{IP: net.IPv4(192, 168, 0, 0), Mask: net.CIDRMask(16, 32)}, // 192.168.0.0/16 (私有)
{IP: net.IPv4(198, 18, 0, 0), Mask: net.CIDRMask(15, 32)}, // 198.18.0.0/15 (基准测试)
{IP: net.IPv4(198, 51, 100, 0), Mask: net.CIDRMask(24, 32)}, // 198.51.100.0/24 (TEST-NET-2)
{IP: net.IPv4(203, 0, 113, 0), Mask: net.CIDRMask(24, 32)}, // 203.0.113.0/24 (TEST-NET-3)
{IP: net.IPv4(224, 0, 0, 0), Mask: net.CIDRMask(4, 32)}, // 224.0.0.0/4 (组播)
{IP: net.IPv4(240, 0, 0, 0), Mask: net.CIDRMask(4, 32)}, // 240.0.0.0/4 (保留)
{IP: net.IPv4(0, 0, 0, 0), Mask: net.CIDRMask(8, 32)}, // 0.0.0.0/8 ("This network" / 未指定)
{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}, // 10.0.0.0/8 (私有)
{IP: net.IPv4(100, 64, 0, 0), Mask: net.CIDRMask(10, 32)}, // 100.64.0.0/10 (运营商级 NAT / CGNAT)
{IP: net.IPv4(127, 0, 0, 0), Mask: net.CIDRMask(8, 32)}, // 127.0.0.0/8 (回环)
{IP: net.IPv4(169, 254, 0, 0), Mask: net.CIDRMask(16, 32)}, // 169.254.0.0/16 (链路本地)
{IP: net.IPv4(172, 16, 0, 0), Mask: net.CIDRMask(12, 32)}, // 172.16.0.0/12 (私有)
{IP: net.IPv4(192, 0, 0, 0), Mask: net.CIDRMask(24, 32)}, // 192.0.0.0/24 (IETF 协议分配)
{IP: net.IPv4(192, 0, 2, 0), Mask: net.CIDRMask(24, 32)}, // 192.0.2.0/24 (TEST-NET-1)
{IP: net.IPv4(192, 168, 0, 0), Mask: net.CIDRMask(16, 32)}, // 192.168.0.0/16 (私有)
{IP: net.IPv4(198, 18, 0, 0), Mask: net.CIDRMask(15, 32)}, // 198.18.0.0/15 (基准测试)
{IP: net.IPv4(198, 51, 100, 0), Mask: net.CIDRMask(24, 32)}, // 198.51.100.0/24 (TEST-NET-2)
{IP: net.IPv4(203, 0, 113, 0), Mask: net.CIDRMask(24, 32)}, // 203.0.113.0/24 (TEST-NET-3)
{IP: net.IPv4(224, 0, 0, 0), Mask: net.CIDRMask(4, 32)}, // 224.0.0.0/4 (组播)
{IP: net.IPv4(240, 0, 0, 0), Mask: net.CIDRMask(4, 32)}, // 240.0.0.0/4 (保留)
{IP: net.IPv4(255, 255, 255, 255), Mask: net.CIDRMask(32, 32)}, // 255.255.255.255/32 (受限广播)
}
@@ -248,6 +266,63 @@ func (p *SSRFProtection) IsIPAccessAllowed(ip net.IP) bool {
return !listed
}
func (p *SSRFProtection) ipAccessError(host string, ip net.IP) error {
if host != "" {
if isPrivateIP(ip) && !p.AllowPrivateIp {
return fmt.Errorf("private IP address not allowed: %s resolves to %s", host, ip.String())
}
if p.IpFilterMode {
return fmt.Errorf("ip not in whitelist: %s resolves to %s", host, ip.String())
}
return fmt.Errorf("ip in blacklist: %s resolves to %s", host, ip.String())
}
if isPrivateIP(ip) && !p.AllowPrivateIp {
return fmt.Errorf("private IP address not allowed: %s", ip.String())
}
if p.IpFilterMode {
return fmt.Errorf("ip not in whitelist: %s", ip.String())
}
return fmt.Errorf("ip in blacklist: %s", ip.String())
}
// ValidateNetworkTarget validates the host and port before dialing.
func (p *SSRFProtection) ValidateNetworkTarget(host string, port int) error {
host = strings.TrimSpace(host)
if host == "" {
return fmt.Errorf("invalid host")
}
if port < 1 || port > 65535 {
return fmt.Errorf("invalid port: %d", port)
}
if !p.isAllowedPort(port) {
return fmt.Errorf("port %d is not allowed", port)
}
if ip := net.ParseIP(host); ip != nil {
if !p.IsIPAccessAllowed(ip) {
return p.ipAccessError("", ip)
}
return nil
}
if !p.isDomainAllowed(host) {
if p.DomainFilterMode {
return fmt.Errorf("domain not in whitelist: %s", host)
}
return fmt.Errorf("domain in blacklist: %s", host)
}
return nil
}
// ValidateResolvedIP validates a domain's resolved IP immediately before dialing it.
func (p *SSRFProtection) ValidateResolvedIP(host string, ip net.IP) error {
if !p.IsIPAccessAllowed(ip) {
return p.ipAccessError(host, ip)
}
return nil
}
// ValidateURL 验证URL是否安全
func (p *SSRFProtection) ValidateURL(urlStr string) error {
// 解析URL
@@ -279,34 +354,12 @@ func (p *SSRFProtection) ValidateURL(urlStr string) error {
return fmt.Errorf("invalid port: %s", portStr)
}
if !p.isAllowedPort(port) {
return fmt.Errorf("port %d is not allowed", port)
if err := p.ValidateNetworkTarget(host, port); err != nil {
return err
}
// 如果 host 是 IP则跳过域名检查
if ip := net.ParseIP(host); ip != nil {
if !p.IsIPAccessAllowed(ip) {
if isPrivateIP(ip) {
return fmt.Errorf("private IP address not allowed: %s", ip.String())
}
if p.IpFilterMode {
return fmt.Errorf("ip not in whitelist: %s", ip.String())
}
return fmt.Errorf("ip in blacklist: %s", ip.String())
}
return nil
}
// 先进行域名过滤
if !p.isDomainAllowed(host) {
if p.DomainFilterMode {
return fmt.Errorf("domain not in whitelist: %s", host)
}
return fmt.Errorf("domain in blacklist: %s", host)
}
// 若未启用对域名应用IP过滤,则到此通过
if !p.ApplyIPFilterForDomain {
// 如果 host 是 IP或未启用对域名应用 IP 过滤,则到此通过。
if net.ParseIP(host) != nil || !p.ApplyIPFilterForDomain {
return nil
}
@@ -316,14 +369,8 @@ func (p *SSRFProtection) ValidateURL(urlStr string) error {
return fmt.Errorf("DNS resolution failed for %s: %v", host, err)
}
for _, ip := range ips {
if !p.IsIPAccessAllowed(ip) {
if isPrivateIP(ip) && !p.AllowPrivateIp {
return fmt.Errorf("private IP address not allowed: %s resolves to %s", host, ip.String())
}
if p.IpFilterMode {
return fmt.Errorf("ip not in whitelist: %s resolves to %s", host, ip.String())
}
return fmt.Errorf("ip in blacklist: %s resolves to %s", host, ip.String())
if err := p.ValidateResolvedIP(host, ip); err != nil {
return err
}
}
return nil
@@ -336,20 +383,9 @@ func ValidateURLWithFetchSetting(urlStr string, enableSSRFProtection, allowPriva
return nil
}
// 解析端口范围配置
allowedPortInts, err := parsePortRanges(allowedPorts)
protection, err := NewSSRFProtectionFromFetchSetting(allowPrivateIp, domainFilterMode, ipFilterMode, domainList, ipList, allowedPorts, applyIPFilterForDomain)
if err != nil {
return fmt.Errorf("request reject - invalid port configuration: %v", err)
}
protection := &SSRFProtection{
AllowPrivateIp: allowPrivateIp,
DomainFilterMode: domainFilterMode,
DomainList: domainList,
IpFilterMode: ipFilterMode,
IpList: ipList,
AllowedPorts: allowedPortInts,
ApplyIPFilterForDomain: applyIPFilterForDomain,
return err
}
return protection.ValidateURL(urlStr)
}
+59
View File
@@ -0,0 +1,59 @@
package common
import (
"net"
"testing"
"github.com/stretchr/testify/require"
)
func TestSSRFProtectionRejectsLiteralPrivateAndReservedIPs(t *testing.T) {
protection := &SSRFProtection{
AllowPrivateIp: false,
DomainFilterMode: false,
IpFilterMode: false,
}
tests := []string{
"127.0.0.1",
"10.0.0.1",
"169.254.169.254",
"fc00::1",
"::ffff:127.0.0.1",
}
for _, host := range tests {
t.Run(host, func(t *testing.T) {
require.Error(t, protection.ValidateNetworkTarget(host, 80))
})
}
}
func TestSSRFProtectionAllowsPrivateIPWhenExplicitlyEnabled(t *testing.T) {
protection := &SSRFProtection{
AllowPrivateIp: true,
DomainFilterMode: false,
IpFilterMode: false,
}
require.NoError(t, protection.ValidateNetworkTarget("10.0.0.1", 80))
}
func TestSSRFProtectionRejectsResolvedPrivateIP(t *testing.T) {
protection := &SSRFProtection{
AllowPrivateIp: false,
DomainFilterMode: false,
IpFilterMode: false,
ApplyIPFilterForDomain: true,
}
require.NoError(t, protection.ValidateNetworkTarget("example.com", 80))
require.Error(t, protection.ValidateResolvedIP("example.com", net.ParseIP("169.254.169.254")))
}
func TestNewSSRFProtectionFromFetchSettingParsesPortRanges(t *testing.T) {
protection, err := NewSSRFProtectionFromFetchSetting(false, false, false, nil, nil, []string{"80", "8000-8001"}, true)
require.NoError(t, err)
require.NoError(t, protection.ValidateNetworkTarget("example.com", 8001))
require.Error(t, protection.ValidateNetworkTarget("example.com", 9000))
}