diff --git a/mtglib/network/default_test.go b/mtglib/network/default_test.go index 203a2bb..95048cf 100644 --- a/mtglib/network/default_test.go +++ b/mtglib/network/default_test.go @@ -10,6 +10,7 @@ import ( ) type DefaultDialerTestSuite struct { + suite.Suite HTTPServerTestSuite d network.Dialer @@ -65,17 +66,12 @@ func (suite *DefaultDialerTestSuite) TestConnectOk() { } func (suite *DefaultDialerTestSuite) TestHTTPRequest() { - httpClient := http.Client{ - Transport: &http.Transport{ - DialContext: suite.d.DialContext, - }, - } + httpClient := suite.MakeHTTPClient(suite.d) - resp, err := httpClient.Get(suite.httpServer.URL + "/get") + resp, err := httpClient.Get(suite.MakeURL("/get")) suite.NoError(err) - - resp.Body.Close() + suite.Equal(http.StatusOK, resp.StatusCode) } func TestDefaultDialer(t *testing.T) { diff --git a/mtglib/network/init_test.go b/mtglib/network/init_test.go index 8f8aeff..5ac8cdd 100644 --- a/mtglib/network/init_test.go +++ b/mtglib/network/init_test.go @@ -1,16 +1,77 @@ package network_test import ( + "context" + "net" + "net/http" "net/http/httptest" + "net/url" "strings" + "time" + "github.com/9seconds/mtg/v2/mtglib/network" + socks5 "github.com/armon/go-socks5" "github.com/mccutchen/go-httpbin/httpbin" - "github.com/stretchr/testify/suite" + "github.com/stretchr/testify/mock" ) -type HTTPServerTestSuite struct { - suite.Suite +type ConnMock struct { + mock.Mock +} +func (c *ConnMock) Read(b []byte) (int, error) { + args := c.Called(b) + + return args.Int(0), args.Error(1) +} + +func (c *ConnMock) Write(b []byte) (int, error) { + args := c.Called(b) + + return args.Int(0), args.Error(1) +} + +func (c *ConnMock) Close() error { + return c.Called().Error(0) +} + +func (c *ConnMock) LocalAddr() net.Addr { + return c.Called().Get(0).(net.Addr) +} + +func (c *ConnMock) RemoteAddr() net.Addr { + return c.Called().Get(0).(net.Addr) +} + +func (c *ConnMock) SetDeadline(t time.Time) error { + return c.Called(t).Error(0) +} + +func (c *ConnMock) SetReadDeadline(t time.Time) error { + return c.Called(t).Error(0) +} + +func (c *ConnMock) SetWriteDeadline(t time.Time) error { + return c.Called(t).Error(0) +} + +type DialerMock struct { + mock.Mock +} + +func (d *DialerMock) Dial(network, address string) (net.Conn, error) { + args := d.Called(network, address) + + return args.Get(0).(net.Conn), args.Error(1) +} + +func (d *DialerMock) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + args := d.Called(ctx, network, address) + + return args.Get(0).(net.Conn), args.Error(1) +} + +type HTTPServerTestSuite struct { httpServer *httptest.Server } @@ -25,3 +86,43 @@ func (suite *HTTPServerTestSuite) TearDownSuite() { func (suite *HTTPServerTestSuite) HTTPServerAddress() string { return strings.TrimPrefix(suite.httpServer.URL, "http://") } + +func (suite *HTTPServerTestSuite) MakeURL(path string) string { + return suite.httpServer.URL + path +} + +func (suite *HTTPServerTestSuite) MakeHTTPClient(dialer network.Dialer) *http.Client { + return &http.Client{ + Transport: &http.Transport{ + DialContext: dialer.DialContext, + }, + } +} + +type Socks5ServerTestSuite struct { + socks5Listener net.Listener + socks5Server *socks5.Server +} + +func (suite *Socks5ServerTestSuite) SetupSuite() { + suite.socks5Listener, _ = net.Listen("tcp", "127.0.0.1:0") + suite.socks5Server, _ = socks5.New(&socks5.Config{ + Credentials: socks5.StaticCredentials{ + "user": "password", + }, + }) + + go suite.socks5Server.Serve(suite.socks5Listener) +} + +func (suite *Socks5ServerTestSuite) TearDownSuite() { + suite.socks5Listener.Close() +} + +func (suite *Socks5ServerTestSuite) MakeSocks5URL(user, password string) *url.URL { + return &url.URL{ + Scheme: "socks5", + User: url.UserPassword(user, password), + Host: suite.socks5Listener.Addr().String(), + } +} diff --git a/mtglib/network/load_balanced_socks5_test.go b/mtglib/network/load_balanced_socks5_test.go new file mode 100644 index 0000000..129597b --- /dev/null +++ b/mtglib/network/load_balanced_socks5_test.go @@ -0,0 +1,88 @@ +package network_test + +import ( + "errors" + "io" + "net" + "net/http" + "net/url" + "testing" + + "github.com/9seconds/mtg/v2/mtglib/network" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/suite" +) + +type LoadBalancedSocks5TestSuite struct { + suite.Suite + HTTPServerTestSuite + Socks5ServerTestSuite + + httpClient *http.Client +} + +func (suite *LoadBalancedSocks5TestSuite) SetupSuite() { + suite.HTTPServerTestSuite.SetupSuite() + suite.Socks5ServerTestSuite.SetupSuite() +} + +func (suite *LoadBalancedSocks5TestSuite) SetupTest() { + baseDialer, _ := network.NewDefaultDialer(0, 0) + + lbDialer, err := network.NewLoadBalancedSocks5Dialer(baseDialer, []*url.URL{ + suite.MakeSocks5URL("user", "password"), + suite.MakeSocks5URL("user2", "password"), + }) + + suite.NoError(err) + + suite.httpClient = suite.MakeHTTPClient(lbDialer) +} + +func (suite *LoadBalancedSocks5TestSuite) TearDownSuite() { + suite.Socks5ServerTestSuite.SetupSuite() + suite.HTTPServerTestSuite.SetupSuite() +} + +func (suite *LoadBalancedSocks5TestSuite) TestIncorrectURL() { + _, err := network.NewLoadBalancedSocks5Dialer(&DialerMock{}, []*url.URL{ + {Scheme: "http"}, + }) + + suite.Error(err) +} + +func (suite *LoadBalancedSocks5TestSuite) TestCannotDial() { + baseDialer := &DialerMock{} + + baseDialer.On("DialContext", mock.Anything, "tcp", "127.0.0.1:1080"). + Times(network.ProxyDialerOpenThreshold). + Return(&net.TCPConn{}, io.EOF) + baseDialer.On("DialContext", mock.Anything, "tcp", "127.0.0.2:1080"). + Times(network.ProxyDialerOpenThreshold). + Return(&net.TCPConn{}, io.EOF) + + lbDialer, err := network.NewLoadBalancedSocks5Dialer(baseDialer, []*url.URL{ + {Scheme: "socks5", User: url.UserPassword("user", "password"), Host: "127.0.0.1:1080"}, + {Scheme: "socks5", User: url.UserPassword("user", "password"), Host: "127.0.0.2:1080"}, + }) + suite.NoError(err) + + for i := 0; i < network.ProxyDialerOpenThreshold*2; i++ { + _, err = lbDialer.Dial("tcp", "127.1.1.1:80") + suite.True(errors.Is(err, network.ErrCannotDialWithAllProxies)) + } + + baseDialer.AssertExpectations(suite.T()) +} + +func (suite *LoadBalancedSocks5TestSuite) TestDialOk() { + resp, err := suite.httpClient.Get(suite.MakeURL("/get")) + + suite.NoError(err) + suite.Equal(http.StatusOK, resp.StatusCode) +} + +func TestLoadBalancedSocks5(t *testing.T) { + suite.Run(t, &LoadBalancedSocks5TestSuite{}) +} diff --git a/mtglib/network/socks5_test.go b/mtglib/network/socks5_test.go index f419a69..5499af4 100644 --- a/mtglib/network/socks5_test.go +++ b/mtglib/network/socks5_test.go @@ -1,84 +1,52 @@ package network_test import ( - "net" "net/http" - "net/url" "testing" "github.com/9seconds/mtg/v2/mtglib/network" - socks5 "github.com/armon/go-socks5" "github.com/stretchr/testify/suite" ) type Socks5TestSuite struct { + suite.Suite HTTPServerTestSuite + Socks5ServerTestSuite - baseDialer network.Dialer - socksListener net.Listener - socksProxy *socks5.Server + d network.Dialer } func (suite *Socks5TestSuite) SetupSuite() { suite.HTTPServerTestSuite.SetupSuite() + suite.Socks5ServerTestSuite.SetupSuite() - socksConf := socks5.Config{ - Credentials: socks5.StaticCredentials{ - "user": "password", - }, - } - - suite.socksProxy, _ = socks5.New(&socksConf) - suite.socksListener, _ = net.Listen("tcp", "127.0.0.1:0") - suite.baseDialer, _ = network.NewDefaultDialer(0, 0) - - go suite.socksProxy.Serve(suite.socksListener) + suite.d, _ = network.NewDefaultDialer(0, 0) } func (suite *Socks5TestSuite) TearDownSuite() { - suite.socksListener.Close() - + suite.Socks5ServerTestSuite.TearDownSuite() suite.HTTPServerTestSuite.TearDownSuite() } func (suite *Socks5TestSuite) TestRequestFailed() { - proxyURL := &url.URL{ - Scheme: "socks5", - User: url.UserPassword("user2", "password"), - Host: suite.socksListener.Addr().String(), - } - dialer, _ := network.NewSocks5Dialer(suite.baseDialer, proxyURL) + proxyURL := suite.MakeSocks5URL("user2", "password") + dialer, _ := network.NewSocks5Dialer(suite.d, proxyURL) + httpClient := suite.MakeHTTPClient(dialer) - httpClient := http.Client{ - Transport: &http.Transport{ - DialContext: dialer.DialContext, - }, - } - - _, err := httpClient.Get(suite.httpServer.URL + "/get") + _, err := httpClient.Get(suite.MakeURL("/get")) suite.Error(err) } func (suite *Socks5TestSuite) TestRequestOk() { - proxyURL := &url.URL{ - Scheme: "socks5", - User: url.UserPassword("user", "password"), - Host: suite.socksListener.Addr().String(), - } - dialer, _ := network.NewSocks5Dialer(suite.baseDialer, proxyURL) + proxyURL := suite.MakeSocks5URL("user", "password") + dialer, _ := network.NewSocks5Dialer(suite.d, proxyURL) + httpClient := suite.MakeHTTPClient(dialer) - httpClient := http.Client{ - Transport: &http.Transport{ - DialContext: dialer.DialContext, - }, - } - - resp, err := httpClient.Get(suite.httpServer.URL + "/get") + resp, err := httpClient.Get(suite.MakeURL("/get")) suite.NoError(err) - - resp.Body.Close() + suite.Equal(http.StatusOK, resp.StatusCode) } func TestSocks5TestSuite(t *testing.T) {