feat(ssrf): implement SSRF protection in HTTP clients and validation functions
This commit is contained in:
+97
-61
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -68,11 +68,16 @@ func VideoProxy(c *gin.Context) {
|
||||
|
||||
var videoURL string
|
||||
proxy := channel.GetSetting().Proxy
|
||||
client, err := service.GetHttpClientWithProxy(proxy)
|
||||
if err != nil {
|
||||
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to create proxy client for task %s: %s", taskID, err.Error()))
|
||||
videoProxyError(c, http.StatusInternalServerError, "server_error", "Failed to create proxy client")
|
||||
return
|
||||
client := service.GetSSRFProtectedHTTPClient()
|
||||
if proxy != "" {
|
||||
// 渠道代理路径的连接由代理侧建立,无法做拨号时逐 IP 校验,
|
||||
// 因此后面对 videoURL 保留请求前的一次性 SSRF 校验。
|
||||
client, err = service.GetHttpClientWithProxy(proxy)
|
||||
if err != nil {
|
||||
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to create proxy client for task %s: %s", taskID, err.Error()))
|
||||
videoProxyError(c, http.StatusInternalServerError, "server_error", "Failed to create proxy client")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
|
||||
@@ -129,10 +134,16 @@ func VideoProxy(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
if err := common.ValidateURLWithFetchSetting(videoURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
|
||||
logger.LogError(c.Request.Context(), fmt.Sprintf("Video URL blocked for task %s: %v", taskID, err))
|
||||
videoProxyError(c, http.StatusForbidden, "server_error", fmt.Sprintf("request blocked: %v", err))
|
||||
var validateErr error
|
||||
if proxy == "" {
|
||||
validateErr = service.ValidateSSRFProtectedFetchURL(videoURL)
|
||||
} else {
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
validateErr = common.ValidateURLWithFetchSetting(videoURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain)
|
||||
}
|
||||
if validateErr != nil {
|
||||
logger.LogError(c.Request.Context(), fmt.Sprintf("Video URL blocked for task %s: %v", taskID, validateErr))
|
||||
videoProxyError(c, http.StatusForbidden, "server_error", fmt.Sprintf("request blocked: %v", validateErr))
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -36,8 +36,9 @@ func RelayMidjourneyImage(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
var httpClient *http.Client
|
||||
var proxy string
|
||||
if channel, err := model.CacheGetChannel(midjourneyTask.ChannelId); err == nil {
|
||||
proxy := channel.GetSetting().Proxy
|
||||
proxy = channel.GetSetting().Proxy
|
||||
if proxy != "" {
|
||||
if httpClient, err = service.NewProxyHttpClient(proxy); err != nil {
|
||||
c.JSON(400, gin.H{
|
||||
@@ -48,12 +49,20 @@ func RelayMidjourneyImage(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
if httpClient == nil {
|
||||
httpClient = service.GetHttpClient()
|
||||
httpClient = service.GetSSRFProtectedHTTPClient()
|
||||
}
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
if err := common.ValidateURLWithFetchSetting(midjourneyTask.ImageUrl, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
|
||||
var validateErr error
|
||||
if proxy == "" {
|
||||
validateErr = service.ValidateSSRFProtectedFetchURL(midjourneyTask.ImageUrl)
|
||||
} else {
|
||||
// 渠道代理路径的连接由代理侧建立,无法做拨号时逐 IP 校验,
|
||||
// 因此保留请求前的一次性 SSRF 校验。
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
validateErr = common.ValidateURLWithFetchSetting(midjourneyTask.ImageUrl, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain)
|
||||
}
|
||||
if validateErr != nil {
|
||||
c.JSON(http.StatusForbidden, gin.H{
|
||||
"error": fmt.Sprintf("request blocked: %v", err),
|
||||
"error": fmt.Sprintf("request blocked: %v", validateErr),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
+3
-4
@@ -41,7 +41,7 @@ func DoWorkerRequest(req *WorkerRequest) (*http.Response, error) {
|
||||
}
|
||||
|
||||
// 序列化worker请求数据
|
||||
workerPayload, err := json.Marshal(req)
|
||||
workerPayload, err := common.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal worker payload: %v", err)
|
||||
}
|
||||
@@ -59,12 +59,11 @@ func DoDownloadRequest(originUrl string, reason ...string) (resp *http.Response,
|
||||
return DoWorkerRequest(req)
|
||||
} else {
|
||||
// SSRF防护:验证请求URL(非Worker模式)
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
if err := common.ValidateURLWithFetchSetting(originUrl, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
|
||||
if err := ValidateSSRFProtectedFetchURL(originUrl); err != nil {
|
||||
return nil, fmt.Errorf("request reject: %v", err)
|
||||
}
|
||||
|
||||
common.SysLog(fmt.Sprintf("downloading from origin: %s, reason: %s", common.MaskSensitiveInfo(originUrl), strings.Join(reason, ", ")))
|
||||
return GetHttpClient().Get(originUrl)
|
||||
return GetSSRFProtectedHTTPClient().Get(originUrl)
|
||||
}
|
||||
}
|
||||
|
||||
+42
-5
@@ -16,15 +16,15 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
httpClient *http.Client
|
||||
proxyClientLock sync.Mutex
|
||||
proxyClients = make(map[string]*http.Client)
|
||||
httpClient *http.Client
|
||||
ssrfProtectedHTTPClient *http.Client
|
||||
proxyClientLock sync.Mutex
|
||||
proxyClients = make(map[string]*http.Client)
|
||||
)
|
||||
|
||||
func checkRedirect(req *http.Request, via []*http.Request) error {
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
urlStr := req.URL.String()
|
||||
if err := common.ValidateURLWithFetchSetting(urlStr, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
|
||||
if err := validateURLWithCurrentFetchSetting(urlStr, true); err != nil {
|
||||
return fmt.Errorf("redirect to %s blocked: %v", urlStr, err)
|
||||
}
|
||||
if len(via) >= 10 {
|
||||
@@ -33,6 +33,26 @@ func checkRedirect(req *http.Request, via []*http.Request) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkProtectedFetchRedirect(req *http.Request, via []*http.Request) error {
|
||||
urlStr := req.URL.String()
|
||||
if err := ValidateSSRFProtectedFetchURL(urlStr); err != nil {
|
||||
return fmt.Errorf("redirect to %s blocked: %v", urlStr, err)
|
||||
}
|
||||
if len(via) >= 10 {
|
||||
return fmt.Errorf("stopped after 10 redirects")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateURLWithCurrentFetchSetting(urlStr string, applyDomainIPFilter bool) error {
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
return common.ValidateURLWithFetchSetting(urlStr, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, applyDomainIPFilter && fetchSetting.ApplyIPFilterForDomain)
|
||||
}
|
||||
|
||||
func ValidateSSRFProtectedFetchURL(urlStr string) error {
|
||||
return validateURLWithCurrentFetchSetting(urlStr, true)
|
||||
}
|
||||
|
||||
func InitHttpClient() {
|
||||
transport := &http.Transport{
|
||||
MaxIdleConns: common.RelayMaxIdleConns,
|
||||
@@ -57,12 +77,29 @@ func InitHttpClient() {
|
||||
CheckRedirect: checkRedirect,
|
||||
}
|
||||
}
|
||||
ssrfProtectedHTTPClient = newProtectedFetchHTTPClient()
|
||||
}
|
||||
|
||||
// GetHttpClient returns the general outbound client used by relay/provider
|
||||
// integrations. Do not attach the SSRF-protected dialer here: provider base URLs
|
||||
// are root/operator-managed deployment targets, not arbitrary user-controlled
|
||||
// input, and may legitimately point at private networks, private-link endpoints,
|
||||
// self-hosted services, or local proxies. Code paths that fetch arbitrary
|
||||
// user-controlled URLs must use GetSSRFProtectedHTTPClient or
|
||||
// ValidateSSRFProtectedFetchURL instead.
|
||||
func GetHttpClient() *http.Client {
|
||||
return httpClient
|
||||
}
|
||||
|
||||
// GetSSRFProtectedHTTPClient 返回带拨号时 SSRF 校验的客户端。
|
||||
// ssrfProtectedHTTPClient 由 InitHttpClient 在启动时初始化,运行期只读。
|
||||
func GetSSRFProtectedHTTPClient() *http.Client {
|
||||
if fetchSetting := system_setting.GetFetchSetting(); fetchSetting != nil && !fetchSetting.EnableSSRFProtection {
|
||||
return GetHttpClient()
|
||||
}
|
||||
return ssrfProtectedHTTPClient
|
||||
}
|
||||
|
||||
// GetHttpClientWithProxy returns the default client or a proxy-enabled one when proxyURL is provided.
|
||||
func GetHttpClientWithProxy(proxyURL string) (*http.Client, error) {
|
||||
if proxyURL == "" {
|
||||
|
||||
@@ -0,0 +1,239 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/QuantumNous/new-api/common"
|
||||
"github.com/QuantumNous/new-api/setting/system_setting"
|
||||
)
|
||||
|
||||
type ssrfResolver interface {
|
||||
LookupIPAddr(ctx context.Context, host string) ([]net.IPAddr, error)
|
||||
}
|
||||
|
||||
type protectedFetchDialer struct {
|
||||
resolver ssrfResolver
|
||||
dialContext func(ctx context.Context, network, address string) (net.Conn, error)
|
||||
getProtection func() (*common.SSRFProtection, bool, error)
|
||||
}
|
||||
|
||||
type ssrfProtectedRoundTripper struct {
|
||||
resolver ssrfResolver
|
||||
dialContext func(ctx context.Context, network, address string) (net.Conn, error)
|
||||
getProtection func() (*common.SSRFProtection, bool, error)
|
||||
proxy func(*http.Request) (*url.URL, error)
|
||||
|
||||
mutex sync.Mutex
|
||||
transports map[string]*http.Transport
|
||||
}
|
||||
|
||||
func currentFetchProtection() (*common.SSRFProtection, bool, error) {
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
if !fetchSetting.EnableSSRFProtection {
|
||||
return nil, false, nil
|
||||
}
|
||||
|
||||
protection, err := common.NewSSRFProtectionFromFetchSetting(
|
||||
fetchSetting.AllowPrivateIp,
|
||||
fetchSetting.DomainFilterMode,
|
||||
fetchSetting.IpFilterMode,
|
||||
fetchSetting.DomainList,
|
||||
fetchSetting.IpList,
|
||||
fetchSetting.AllowedPorts,
|
||||
fetchSetting.ApplyIPFilterForDomain,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, true, err
|
||||
}
|
||||
return protection, true, nil
|
||||
}
|
||||
|
||||
func newProtectedFetchHTTPClient() *http.Client {
|
||||
return newProtectedFetchHTTPClientWithDialer(nil, nil, nil)
|
||||
}
|
||||
|
||||
func newProtectedFetchHTTPClientWithDialer(resolver ssrfResolver, dialContext func(ctx context.Context, network, address string) (net.Conn, error), getProtection func() (*common.SSRFProtection, bool, error)) *http.Client {
|
||||
return newProtectedFetchHTTPClientWithProxy(resolver, dialContext, getProtection, http.ProxyFromEnvironment)
|
||||
}
|
||||
|
||||
func newProtectedFetchHTTPClientWithProxy(resolver ssrfResolver, dialContext func(ctx context.Context, network, address string) (net.Conn, error), getProtection func() (*common.SSRFProtection, bool, error), proxy func(*http.Request) (*url.URL, error)) *http.Client {
|
||||
if resolver == nil {
|
||||
resolver = net.DefaultResolver
|
||||
}
|
||||
if dialContext == nil {
|
||||
netDialer := &net.Dialer{
|
||||
Timeout: 30 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
dialContext = netDialer.DialContext
|
||||
}
|
||||
if getProtection == nil {
|
||||
getProtection = currentFetchProtection
|
||||
}
|
||||
if proxy == nil {
|
||||
proxy = http.ProxyFromEnvironment
|
||||
}
|
||||
|
||||
client := &http.Client{
|
||||
Transport: &ssrfProtectedRoundTripper{
|
||||
resolver: resolver,
|
||||
dialContext: dialContext,
|
||||
getProtection: getProtection,
|
||||
proxy: proxy,
|
||||
transports: make(map[string]*http.Transport),
|
||||
},
|
||||
CheckRedirect: checkProtectedFetchRedirect,
|
||||
}
|
||||
if common.RelayTimeout != 0 {
|
||||
client.Timeout = time.Duration(common.RelayTimeout) * time.Second
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
func (t *ssrfProtectedRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if req == nil || req.URL == nil {
|
||||
return nil, fmt.Errorf("invalid request")
|
||||
}
|
||||
if err := ValidateSSRFProtectedFetchURL(req.URL.String()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
proxyURL, err := t.proxy(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return t.transportFor(proxyURL).RoundTrip(req)
|
||||
}
|
||||
|
||||
func (t *ssrfProtectedRoundTripper) CloseIdleConnections() {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
|
||||
for _, transport := range t.transports {
|
||||
transport.CloseIdleConnections()
|
||||
}
|
||||
}
|
||||
|
||||
func (t *ssrfProtectedRoundTripper) transportFor(proxyURL *url.URL) *http.Transport {
|
||||
// 只按代理地址分组:代理来自环境变量,取值有限,map 有界;
|
||||
// 目标 origin 是用户可控输入,不能作为缓存 key。
|
||||
key := "direct"
|
||||
if proxyURL != nil {
|
||||
key = proxyURL.String()
|
||||
}
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
|
||||
if transport, ok := t.transports[key]; ok {
|
||||
return transport
|
||||
}
|
||||
|
||||
transport := t.newTransport(proxyURL)
|
||||
t.transports[key] = transport
|
||||
return transport
|
||||
}
|
||||
|
||||
func (t *ssrfProtectedRoundTripper) newTransport(proxyURL *url.URL) *http.Transport {
|
||||
dialContext := t.dialContext
|
||||
proxyFunc := http.ProxyURL(proxyURL)
|
||||
if proxyURL == nil {
|
||||
protectedDialer := &protectedFetchDialer{
|
||||
resolver: t.resolver,
|
||||
dialContext: t.dialContext,
|
||||
getProtection: t.getProtection,
|
||||
}
|
||||
dialContext = protectedDialer.DialContext
|
||||
proxyFunc = nil
|
||||
}
|
||||
|
||||
transport := &http.Transport{
|
||||
MaxIdleConns: common.RelayMaxIdleConns,
|
||||
MaxIdleConnsPerHost: common.RelayMaxIdleConnsPerHost,
|
||||
IdleConnTimeout: time.Duration(common.RelayIdleConnTimeout) * time.Second,
|
||||
ForceAttemptHTTP2: true,
|
||||
Proxy: proxyFunc,
|
||||
DialContext: dialContext,
|
||||
}
|
||||
if common.TLSInsecureSkipVerify {
|
||||
transport.TLSClientConfig = common.InsecureTLSConfig
|
||||
}
|
||||
return transport
|
||||
}
|
||||
|
||||
func (d *protectedFetchDialer) DialContext(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
protection, enabled, err := d.getProtection()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !enabled {
|
||||
return d.dialContext(ctx, network, addr)
|
||||
}
|
||||
|
||||
host, portText, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid dial address %s: %w", addr, err)
|
||||
}
|
||||
port, err := strconv.Atoi(portText)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid port: %s", portText)
|
||||
}
|
||||
if err := protection.ValidateNetworkTarget(host, port); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return d.dialContext(ctx, network, net.JoinHostPort(ip.String(), portText))
|
||||
}
|
||||
if !protection.ApplyIPFilterForDomain {
|
||||
return d.dialContext(ctx, network, addr)
|
||||
}
|
||||
|
||||
resolved, err := d.resolver.LookupIPAddr(ctx, host)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("DNS resolution failed for %s: %v", host, err)
|
||||
}
|
||||
|
||||
var candidateIPs []net.IP
|
||||
for _, ipAddr := range resolved {
|
||||
ip := ipAddr.IP
|
||||
if ip == nil || !networkAllowsIP(network, ip) {
|
||||
continue
|
||||
}
|
||||
if err := protection.ValidateResolvedIP(host, ip); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
candidateIPs = append(candidateIPs, ip)
|
||||
}
|
||||
|
||||
var lastDialErr error
|
||||
for _, ip := range candidateIPs {
|
||||
conn, err := d.dialContext(ctx, network, net.JoinHostPort(ip.String(), portText))
|
||||
if err == nil {
|
||||
return conn, nil
|
||||
}
|
||||
lastDialErr = err
|
||||
}
|
||||
|
||||
if lastDialErr != nil {
|
||||
return nil, lastDialErr
|
||||
}
|
||||
return nil, fmt.Errorf("DNS resolution for %s returned no usable IP addresses", host)
|
||||
}
|
||||
|
||||
func networkAllowsIP(network string, ip net.IP) bool {
|
||||
switch network {
|
||||
case "tcp4":
|
||||
return ip.To4() != nil
|
||||
case "tcp6":
|
||||
return ip.To4() == nil
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
"github.com/QuantumNous/new-api/common"
|
||||
"github.com/QuantumNous/new-api/setting/system_setting"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type staticSSRFResolver map[string][]net.IPAddr
|
||||
|
||||
func (r staticSSRFResolver) LookupIPAddr(ctx context.Context, host string) ([]net.IPAddr, error) {
|
||||
if ips, ok := r[host]; ok {
|
||||
return ips, nil
|
||||
}
|
||||
return nil, fmt.Errorf("unexpected lookup for %s", host)
|
||||
}
|
||||
|
||||
func staticProtection(protection *common.SSRFProtection) func() (*common.SSRFProtection, bool, error) {
|
||||
return func() (*common.SSRFProtection, bool, error) {
|
||||
return protection, true, nil
|
||||
}
|
||||
}
|
||||
|
||||
func testConn(t *testing.T) net.Conn {
|
||||
t.Helper()
|
||||
clientConn, serverConn := net.Pipe()
|
||||
t.Cleanup(func() {
|
||||
clientConn.Close()
|
||||
serverConn.Close()
|
||||
})
|
||||
return clientConn
|
||||
}
|
||||
|
||||
func configureSSRFTestFetchSetting(t *testing.T) {
|
||||
t.Helper()
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
original := *fetchSetting
|
||||
t.Cleanup(func() {
|
||||
*fetchSetting = original
|
||||
})
|
||||
|
||||
fetchSetting.EnableSSRFProtection = true
|
||||
fetchSetting.AllowPrivateIp = false
|
||||
fetchSetting.DomainFilterMode = false
|
||||
fetchSetting.IpFilterMode = false
|
||||
fetchSetting.DomainList = nil
|
||||
fetchSetting.IpList = nil
|
||||
fetchSetting.AllowedPorts = []string{"80", "443"}
|
||||
fetchSetting.ApplyIPFilterForDomain = true
|
||||
}
|
||||
|
||||
func mustParseURL(t *testing.T, rawURL string) *url.URL {
|
||||
t.Helper()
|
||||
parsedURL, err := url.Parse(rawURL)
|
||||
require.NoError(t, err)
|
||||
return parsedURL
|
||||
}
|
||||
|
||||
func TestProtectedFetchDialerRejectsPrivateReboundAddress(t *testing.T) {
|
||||
dialer := &protectedFetchDialer{
|
||||
resolver: staticSSRFResolver{
|
||||
"safe.example": {{IP: net.ParseIP("127.0.0.1")}},
|
||||
},
|
||||
dialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
t.Fatalf("dialContext should not be called for blocked address %s", address)
|
||||
return nil, nil
|
||||
},
|
||||
getProtection: staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: false,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: false,
|
||||
ApplyIPFilterForDomain: true,
|
||||
}),
|
||||
}
|
||||
|
||||
conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:80")
|
||||
|
||||
require.Error(t, err)
|
||||
require.Nil(t, conn)
|
||||
require.Contains(t, err.Error(), "private IP address not allowed")
|
||||
}
|
||||
|
||||
func TestProtectedFetchDialerRejectsMixedResolvedIPs(t *testing.T) {
|
||||
var dialed []string
|
||||
dialer := &protectedFetchDialer{
|
||||
resolver: staticSSRFResolver{
|
||||
"safe.example": {
|
||||
{IP: net.ParseIP("10.0.0.1")},
|
||||
{IP: net.ParseIP("8.8.8.8")},
|
||||
},
|
||||
},
|
||||
dialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
dialed = append(dialed, address)
|
||||
return testConn(t), nil
|
||||
},
|
||||
getProtection: staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: false,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: false,
|
||||
ApplyIPFilterForDomain: true,
|
||||
}),
|
||||
}
|
||||
|
||||
conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:443")
|
||||
require.Error(t, err)
|
||||
require.Nil(t, conn)
|
||||
|
||||
require.Empty(t, dialed)
|
||||
require.Contains(t, err.Error(), "private IP address not allowed")
|
||||
}
|
||||
|
||||
func TestProtectedFetchDialerDialsWhenAllResolvedIPsAllowed(t *testing.T) {
|
||||
var dialed []string
|
||||
dialer := &protectedFetchDialer{
|
||||
resolver: staticSSRFResolver{
|
||||
"safe.example": {
|
||||
{IP: net.ParseIP("8.8.8.8")},
|
||||
{IP: net.ParseIP("1.1.1.1")},
|
||||
},
|
||||
},
|
||||
dialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
dialed = append(dialed, address)
|
||||
return testConn(t), nil
|
||||
},
|
||||
getProtection: staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: false,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: false,
|
||||
ApplyIPFilterForDomain: true,
|
||||
}),
|
||||
}
|
||||
|
||||
conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:443")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, conn)
|
||||
|
||||
require.Equal(t, []string{"8.8.8.8:443"}, dialed)
|
||||
}
|
||||
|
||||
func TestProtectedFetchDialerAllowsPrivateIPWhenWhitelisted(t *testing.T) {
|
||||
var dialed []string
|
||||
dialer := &protectedFetchDialer{
|
||||
resolver: staticSSRFResolver{
|
||||
"internal.example": {{IP: net.ParseIP("10.1.2.3")}},
|
||||
},
|
||||
dialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
dialed = append(dialed, address)
|
||||
return testConn(t), nil
|
||||
},
|
||||
getProtection: staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: true,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: true,
|
||||
IpList: []string{"10.0.0.0/8"},
|
||||
ApplyIPFilterForDomain: true,
|
||||
}),
|
||||
}
|
||||
|
||||
conn, err := dialer.DialContext(context.Background(), "tcp", "internal.example:80")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, conn)
|
||||
|
||||
require.Equal(t, []string{"10.1.2.3:80"}, dialed)
|
||||
}
|
||||
|
||||
func TestProtectedFetchDialerSkipsResolvedIPCheckWhenDisabled(t *testing.T) {
|
||||
var dialed []string
|
||||
dialer := &protectedFetchDialer{
|
||||
resolver: staticSSRFResolver{},
|
||||
dialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
dialed = append(dialed, address)
|
||||
return testConn(t), nil
|
||||
},
|
||||
getProtection: staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: false,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: false,
|
||||
ApplyIPFilterForDomain: false,
|
||||
}),
|
||||
}
|
||||
|
||||
conn, err := dialer.DialContext(context.Background(), "tcp", "safe.example:80")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, conn)
|
||||
|
||||
require.Equal(t, []string{"safe.example:80"}, dialed)
|
||||
}
|
||||
|
||||
func TestGetSSRFProtectedHTTPClientFallsBackToDefaultClientWhenProtectionDisabled(t *testing.T) {
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
originalFetchSetting := *fetchSetting
|
||||
originalHTTPClient := httpClient
|
||||
originalProtectedClient := ssrfProtectedHTTPClient
|
||||
t.Cleanup(func() {
|
||||
*fetchSetting = originalFetchSetting
|
||||
httpClient = originalHTTPClient
|
||||
ssrfProtectedHTTPClient = originalProtectedClient
|
||||
})
|
||||
|
||||
fetchSetting.EnableSSRFProtection = false
|
||||
expected := &http.Client{}
|
||||
httpClient = expected
|
||||
ssrfProtectedHTTPClient = &http.Client{}
|
||||
|
||||
require.Same(t, expected, GetSSRFProtectedHTTPClient())
|
||||
}
|
||||
|
||||
func TestProtectedFetchRoundTripperUsesConfiguredProxy(t *testing.T) {
|
||||
configureSSRFTestFetchSetting(t)
|
||||
proxyURL := mustParseURL(t, "http://127.0.0.1:3128")
|
||||
var dialed []string
|
||||
client := newProtectedFetchHTTPClientWithProxy(
|
||||
staticSSRFResolver{},
|
||||
func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
dialed = append(dialed, address)
|
||||
return nil, errors.New("stop after proxy dial")
|
||||
},
|
||||
staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: false,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: false,
|
||||
ApplyIPFilterForDomain: true,
|
||||
}),
|
||||
func(req *http.Request) (*url.URL, error) {
|
||||
return proxyURL, nil
|
||||
},
|
||||
)
|
||||
req, err := http.NewRequest(http.MethodGet, "http://93.184.216.34/resource", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
require.Error(t, err)
|
||||
require.Nil(t, resp)
|
||||
require.Equal(t, []string{"127.0.0.1:3128"}, dialed)
|
||||
}
|
||||
|
||||
func TestProtectedFetchRoundTripperRejectsPrivateTargetBeforeProxy(t *testing.T) {
|
||||
configureSSRFTestFetchSetting(t)
|
||||
proxyURL := mustParseURL(t, "http://127.0.0.1:3128")
|
||||
var dialed []string
|
||||
client := newProtectedFetchHTTPClientWithProxy(
|
||||
staticSSRFResolver{},
|
||||
func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
dialed = append(dialed, address)
|
||||
return nil, errors.New("proxy should not be dialed")
|
||||
},
|
||||
staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: false,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: false,
|
||||
ApplyIPFilterForDomain: true,
|
||||
}),
|
||||
func(req *http.Request) (*url.URL, error) {
|
||||
return proxyURL, nil
|
||||
},
|
||||
)
|
||||
req, err := http.NewRequest(http.MethodGet, "http://localhost/resource", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
require.Error(t, err)
|
||||
require.Nil(t, resp)
|
||||
require.Contains(t, err.Error(), "private IP address not allowed")
|
||||
require.Empty(t, dialed)
|
||||
}
|
||||
|
||||
func TestProtectedFetchRoundTripperNoProxyUsesProtectedDialer(t *testing.T) {
|
||||
configureSSRFTestFetchSetting(t)
|
||||
var dialed []string
|
||||
client := newProtectedFetchHTTPClientWithProxy(
|
||||
staticSSRFResolver{},
|
||||
func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
dialed = append(dialed, address)
|
||||
return nil, errors.New("unexpected direct dial")
|
||||
},
|
||||
staticProtection(&common.SSRFProtection{
|
||||
AllowPrivateIp: false,
|
||||
DomainFilterMode: false,
|
||||
IpFilterMode: false,
|
||||
ApplyIPFilterForDomain: true,
|
||||
}),
|
||||
func(req *http.Request) (*url.URL, error) {
|
||||
return nil, nil
|
||||
},
|
||||
)
|
||||
req, err := http.NewRequest(http.MethodGet, "http://127.0.0.1/resource", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
require.Error(t, err)
|
||||
require.Nil(t, resp)
|
||||
require.Contains(t, err.Error(), "private IP address not allowed")
|
||||
require.Empty(t, dialed)
|
||||
}
|
||||
|
||||
func TestProtectedFetchRoundTripperReusesTransportPerProxy(t *testing.T) {
|
||||
client := newProtectedFetchHTTPClientWithDialer(nil, nil, nil)
|
||||
roundTripper, ok := client.Transport.(*ssrfProtectedRoundTripper)
|
||||
require.True(t, ok)
|
||||
|
||||
direct := roundTripper.transportFor(nil)
|
||||
directAgain := roundTripper.transportFor(nil)
|
||||
proxied := roundTripper.transportFor(mustParseURL(t, "http://127.0.0.1:3128"))
|
||||
|
||||
require.Same(t, direct, directAgain)
|
||||
require.NotSame(t, direct, proxied)
|
||||
require.True(t, direct.ForceAttemptHTTP2)
|
||||
require.False(t, direct.DisableKeepAlives)
|
||||
}
|
||||
@@ -2,7 +2,6 @@ package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
@@ -154,8 +153,7 @@ func sendBarkNotify(barkURL string, data dto.Notify) error {
|
||||
}
|
||||
} else {
|
||||
// SSRF防护:验证Bark URL(非Worker模式)
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
if err := common.ValidateURLWithFetchSetting(finalURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
|
||||
if err := ValidateSSRFProtectedFetchURL(finalURL); err != nil {
|
||||
return fmt.Errorf("request reject: %v", err)
|
||||
}
|
||||
|
||||
@@ -169,7 +167,7 @@ func sendBarkNotify(barkURL string, data dto.Notify) error {
|
||||
req.Header.Set("User-Agent", "OneAPI-Bark-Notify/1.0")
|
||||
|
||||
// 发送请求
|
||||
client := GetHttpClient()
|
||||
client := GetSSRFProtectedHTTPClient()
|
||||
resp, err = client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to send bark request: %v", err)
|
||||
@@ -215,7 +213,7 @@ func sendGotifyNotify(gotifyUrl string, gotifyToken string, priority int, data d
|
||||
}
|
||||
|
||||
// 序列化为 JSON
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
payloadBytes, err := common.Marshal(payload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal gotify payload: %v", err)
|
||||
}
|
||||
@@ -248,8 +246,7 @@ func sendGotifyNotify(gotifyUrl string, gotifyToken string, priority int, data d
|
||||
}
|
||||
} else {
|
||||
// SSRF防护:验证Gotify URL(非Worker模式)
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
if err := common.ValidateURLWithFetchSetting(finalURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
|
||||
if err := ValidateSSRFProtectedFetchURL(finalURL); err != nil {
|
||||
return fmt.Errorf("request reject: %v", err)
|
||||
}
|
||||
|
||||
@@ -264,7 +261,7 @@ func sendGotifyNotify(gotifyUrl string, gotifyToken string, priority int, data d
|
||||
req.Header.Set("User-Agent", "NewAPI-Gotify-Notify/1.0")
|
||||
|
||||
// 发送请求
|
||||
client := GetHttpClient()
|
||||
client := GetSSRFProtectedHTTPClient()
|
||||
resp, err = client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to send gotify request: %v", err)
|
||||
|
||||
+3
-5
@@ -5,7 +5,6 @@ import (
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"time"
|
||||
@@ -49,7 +48,7 @@ func SendWebhookNotify(webhookURL string, secret string, data dto.Notify) error
|
||||
}
|
||||
|
||||
// 序列化负载
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
payloadBytes, err := common.Marshal(payload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal webhook payload: %v", err)
|
||||
}
|
||||
@@ -89,8 +88,7 @@ func SendWebhookNotify(webhookURL string, secret string, data dto.Notify) error
|
||||
}
|
||||
} else {
|
||||
// SSRF防护:验证Webhook URL(非Worker模式)
|
||||
fetchSetting := system_setting.GetFetchSetting()
|
||||
if err := common.ValidateURLWithFetchSetting(webhookURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
|
||||
if err := ValidateSSRFProtectedFetchURL(webhookURL); err != nil {
|
||||
return fmt.Errorf("request reject: %v", err)
|
||||
}
|
||||
|
||||
@@ -109,7 +107,7 @@ func SendWebhookNotify(webhookURL string, secret string, data dto.Notify) error
|
||||
}
|
||||
|
||||
// 发送请求
|
||||
client := GetHttpClient()
|
||||
client := GetSSRFProtectedHTTPClient()
|
||||
resp, err = client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to send webhook request: %v", err)
|
||||
|
||||
Reference in New Issue
Block a user