Files
new-api/service/protected_fetch_client_test.go

318 lines
9.3 KiB
Go

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)
}