Refactor network to a top-level module

This commit is contained in:
9seconds
2021-03-14 21:43:30 +03:00
parent 37a78bd1c3
commit f8ad90c845
25 changed files with 231 additions and 192 deletions
+196
View File
@@ -0,0 +1,196 @@
package network
import (
"context"
"net"
"sync/atomic"
"time"
)
const (
circuitBreakerStateClosed uint32 = iota
circuitBreakerStateHalfOpened
circuitBreakerStateOpened
)
type circuitBreakerDialer struct {
Dialer
stateMutexChan chan bool
halfOpenTimer *time.Timer
failuresCleanupTimer *time.Timer
state uint32
halfOpenAttempts uint32
failuresCount uint32
openThreshold uint32
halfOpenTimeout time.Duration
resetFailuresTimeout time.Duration
}
func (c *circuitBreakerDialer) Dial(network, address string) (net.Conn, error) {
return c.DialContext(context.Background(), network, address)
}
func (c *circuitBreakerDialer) DialContext(ctx context.Context,
network, address string) (net.Conn, error) {
switch atomic.LoadUint32(&c.state) {
case circuitBreakerStateClosed:
return c.doClosed(ctx, network, address)
case circuitBreakerStateHalfOpened:
return c.doHalfOpened(ctx, network, address)
default:
return nil, ErrCircuitBreakerOpened
}
}
func (c *circuitBreakerDialer) doClosed(ctx context.Context,
network, address string) (net.Conn, error) {
conn, err := c.Dialer.DialContext(ctx, network, address)
select {
case <-ctx.Done():
if conn != nil {
conn.Close()
}
return nil, ctx.Err()
case c.stateMutexChan <- true:
defer func() {
<-c.stateMutexChan
}()
}
if err == nil {
c.switchState(circuitBreakerStateClosed)
return conn, err // nolint: wrapcheck
}
c.failuresCount++
if c.state == circuitBreakerStateClosed && c.failuresCount >= c.openThreshold {
c.switchState(circuitBreakerStateOpened)
}
return conn, err // nolint: wrapcheck
}
func (c *circuitBreakerDialer) doHalfOpened(ctx context.Context, network, address string) (net.Conn, error) {
if !atomic.CompareAndSwapUint32(&c.halfOpenAttempts, 0, 1) {
return nil, ErrCircuitBreakerOpened
}
conn, err := c.Dialer.DialContext(ctx, network, address)
select {
case <-ctx.Done():
if conn != nil {
conn.Close()
}
return nil, ctx.Err()
case c.stateMutexChan <- true:
defer func() {
<-c.stateMutexChan
}()
}
if c.state != circuitBreakerStateHalfOpened {
return conn, err // nolint: wrapcheck
}
if err == nil {
c.switchState(circuitBreakerStateClosed)
} else {
c.switchState(circuitBreakerStateOpened)
}
return conn, err // nolint: wrapcheck
}
func (c *circuitBreakerDialer) switchState(state uint32) {
switch state {
case circuitBreakerStateClosed:
c.stopTimer(&c.halfOpenTimer)
c.ensureTimer(&c.failuresCleanupTimer, c.resetFailuresTimeout, c.resetFailures)
case circuitBreakerStateHalfOpened:
c.stopTimer(&c.failuresCleanupTimer)
c.stopTimer(&c.halfOpenTimer)
case circuitBreakerStateOpened:
c.stopTimer(&c.failuresCleanupTimer)
c.ensureTimer(&c.halfOpenTimer, c.halfOpenTimeout, c.tryHalfOpen)
}
c.failuresCount = 0
atomic.StoreUint32(&c.halfOpenAttempts, 0)
atomic.StoreUint32(&c.state, state)
}
func (c *circuitBreakerDialer) resetFailures() {
c.stateMutexChan <- true
defer func() {
<-c.stateMutexChan
}()
c.stopTimer(&c.failuresCleanupTimer)
if c.state == circuitBreakerStateClosed {
c.switchState(circuitBreakerStateClosed)
}
}
func (c *circuitBreakerDialer) tryHalfOpen() {
c.stateMutexChan <- true
defer func() {
<-c.stateMutexChan
}()
if c.state == circuitBreakerStateOpened {
c.switchState(circuitBreakerStateHalfOpened)
}
}
func (c *circuitBreakerDialer) stopTimer(timerRef **time.Timer) {
timer := *timerRef
if timer == nil {
return
}
timer.Stop()
select {
case <-timer.C:
default:
}
*timerRef = nil
}
func (c *circuitBreakerDialer) ensureTimer(timerRef **time.Timer,
timeout time.Duration, callback func()) {
if *timerRef == nil {
*timerRef = time.AfterFunc(timeout, callback)
}
}
func newCircuitBreakerDialer(baseDialer Dialer,
openThreshold uint32, halfOpenTimeout, resetFailuresTimeout time.Duration) Dialer {
cb := &circuitBreakerDialer{
Dialer: baseDialer,
stateMutexChan: make(chan bool, 1),
openThreshold: openThreshold,
halfOpenTimeout: halfOpenTimeout,
resetFailuresTimeout: resetFailuresTimeout,
}
cb.stateMutexChan <- true // to convince race detector we are good
cb.switchState(circuitBreakerStateClosed)
<-cb.stateMutexChan
return cb
}
+139
View File
@@ -0,0 +1,139 @@
package network
import (
"context"
"errors"
"io"
"net"
"sync"
"testing"
"time"
"github.com/9seconds/mtg/v2/testlib"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/suite"
)
type CircuitBreakerTestSuite struct {
suite.Suite
d Dialer
mutex sync.Mutex
ctx context.Context
ctxCancel context.CancelFunc
connMock *testlib.NetConnMock
baseDialerMock *DialerMock
}
func (suite *CircuitBreakerTestSuite) SetupTest() {
suite.mutex = sync.Mutex{}
suite.ctx, suite.ctxCancel = context.WithCancel(context.Background())
suite.baseDialerMock = &DialerMock{}
suite.connMock = &testlib.NetConnMock{}
suite.d = newCircuitBreakerDialer(suite.baseDialerMock,
3, 100*time.Millisecond, 50*time.Millisecond)
}
func (suite *CircuitBreakerTestSuite) TearDownTest() {
suite.ctxCancel()
suite.baseDialerMock.AssertExpectations(suite.T())
suite.connMock.AssertExpectations(suite.T())
}
func (suite *CircuitBreakerTestSuite) TestMultipleRunsOk() {
suite.connMock.On("RemoteAddr").
Times(5).
Return(&net.TCPAddr{
IP: net.ParseIP("127.0.0.1"),
Port: 3128,
})
suite.baseDialerMock.On("DialContext", mock.Anything, "tcp", "127.0.0.1").
Times(5).
Return(suite.connMock, nil)
wg := &sync.WaitGroup{}
wg.Add(5)
go func() {
wg.Wait()
suite.ctxCancel()
}()
for i := 0; i < 5; i++ {
go func() {
defer wg.Done()
conn, err := suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1")
suite.mutex.Lock()
defer suite.mutex.Unlock()
suite.NoError(err)
suite.Equal("127.0.0.1:3128", conn.RemoteAddr().String())
}()
}
suite.Eventually(func() bool {
_, ok := <-suite.ctx.Done()
return !ok
}, time.Second, 10*time.Millisecond)
}
func (suite *CircuitBreakerTestSuite) TestFromClosedToOpen() {
suite.baseDialerMock.On("DialContext", mock.Anything, "tcp", "127.0.0.1").
Times(3).
Return(&net.TCPConn{}, io.EOF)
_, err := suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1")
suite.True(errors.Is(err, io.EOF))
_, err = suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1")
suite.True(errors.Is(err, io.EOF))
_, err = suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1")
suite.True(errors.Is(err, io.EOF))
_, err = suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1")
suite.True(errors.Is(err, ErrCircuitBreakerOpened))
}
func (suite *CircuitBreakerTestSuite) TestHalfOpen() {
suite.baseDialerMock.On("DialContext", mock.Anything, "tcp", "127.0.0.1").
Times(4).
Return(&net.TCPConn{}, io.EOF)
suite.baseDialerMock.On("DialContext", mock.Anything, "tcp", "127.0.0.2").
Twice().
Return(suite.connMock, nil)
suite.connMock.On("RemoteAddr").Return(&net.TCPAddr{
IP: net.ParseIP("10.0.0.10"),
Port: 80,
})
suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1") // nolint: errcheck
suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1") // nolint: errcheck
suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1") // nolint: errcheck
suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1") // nolint: errcheck
time.Sleep(500 * time.Millisecond)
_, err := suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1")
suite.True(errors.Is(err, io.EOF))
_, err = suite.d.DialContext(suite.ctx, "tcp", "127.0.0.1")
suite.True(errors.Is(err, ErrCircuitBreakerOpened))
time.Sleep(500 * time.Millisecond)
conn, err := suite.d.DialContext(suite.ctx, "tcp", "127.0.0.2")
suite.NoError(err)
suite.Equal("10.0.0.10:80", conn.RemoteAddr().String())
_, err = suite.d.DialContext(suite.ctx, "tcp", "127.0.0.2")
suite.NoError(err)
}
func TestCircuitBreaker(t *testing.T) {
t.Parallel()
suite.Run(t, &CircuitBreakerTestSuite{})
}
+86
View File
@@ -0,0 +1,86 @@
package network
import (
"context"
"fmt"
"net"
"time"
"github.com/libp2p/go-reuseport"
)
type defaultDialer struct {
net.Dialer
bufferSize int
}
func (d *defaultDialer) Dial(network, address string) (net.Conn, error) {
return d.DialContext(context.Background(), network, address)
}
func (d *defaultDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
switch network {
case "tcp", "tcp4", "tcp6": // nolint: goconst
default:
return nil, fmt.Errorf("unsupported network %s", network)
}
conn, err := d.Dialer.DialContext(ctx, network, address)
if err != nil {
return nil, fmt.Errorf("cannot dial to %s: %w", address, err)
}
tcpConn := conn.(*net.TCPConn)
if err := tcpConn.SetNoDelay(true); err != nil {
conn.Close()
return nil, fmt.Errorf("cannot set TCP_NO_DELAY: %w", err)
}
if err := tcpConn.SetReadBuffer(d.bufferSize); err != nil {
tcpConn.Close()
return nil, fmt.Errorf("cannot set read buffer size: %w", err)
}
if err := tcpConn.SetWriteBuffer(d.bufferSize); err != nil {
tcpConn.Close()
return nil, fmt.Errorf("cannot set write buffer size: %w", err)
}
if err := tcpConn.SetKeepAlive(true); err != nil {
tcpConn.Close()
return nil, fmt.Errorf("cannot enable keep-alive: %w", err)
}
return tcpConn, nil
}
func NewDefaultDialer(timeout time.Duration, bufferSize int) (Dialer, error) {
switch {
case timeout < 0:
return nil, fmt.Errorf("timeout %v should be positive number", timeout)
case bufferSize < 0:
return nil, fmt.Errorf("buffer size %d should be positive number", bufferSize)
}
if timeout == 0 {
timeout = DefaultTimeout
}
if bufferSize == 0 {
bufferSize = DefaultBufferSize
}
return &defaultDialer{
Dialer: net.Dialer{
Timeout: timeout,
Control: reuseport.Control,
},
bufferSize: bufferSize,
}, nil
}
+77
View File
@@ -0,0 +1,77 @@
package network_test
import (
"context"
"net/http"
"testing"
"github.com/9seconds/mtg/v2/network"
"github.com/stretchr/testify/suite"
)
type DefaultDialerTestSuite struct {
suite.Suite
HTTPServerTestSuite
d network.Dialer
}
func (suite *DefaultDialerTestSuite) SetupSuite() {
suite.HTTPServerTestSuite.SetupSuite()
d, err := network.NewDefaultDialer(0, 0)
suite.NoError(err)
suite.d = d
}
func (suite *DefaultDialerTestSuite) TestNegativeTimeout() {
_, err := network.NewDefaultDialer(-1, 0)
suite.Error(err)
}
func (suite *DefaultDialerTestSuite) TestNegativeBufferSize() {
_, err := network.NewDefaultDialer(0, -1)
suite.Error(err)
}
func (suite *DefaultDialerTestSuite) TestUnsupportedProtocol() {
_, err := suite.d.DialContext(context.Background(),
"udp",
suite.HTTPServerAddress())
suite.Error(err)
}
func (suite *DefaultDialerTestSuite) TestCannotDial() {
_, err := suite.d.DialContext(context.Background(),
"tcp",
suite.HTTPServerAddress()+suite.HTTPServerAddress())
suite.Error(err)
}
func (suite *DefaultDialerTestSuite) TestConnectOk() {
conn, err := suite.d.DialContext(context.Background(),
"tcp",
suite.HTTPServerAddress())
suite.NoError(err)
suite.NotNil(conn)
conn.Close()
}
func (suite *DefaultDialerTestSuite) TestHTTPRequest() {
httpClient := suite.MakeHTTPClient(suite.d)
resp, err := httpClient.Get(suite.MakeURL("/get")) // nolint: noctx
if err == nil {
defer resp.Body.Close()
}
suite.NoError(err)
suite.Equal(http.StatusOK, resp.StatusCode)
}
func TestDefaultDialer(t *testing.T) {
t.Parallel()
suite.Run(t, &DefaultDialerTestSuite{})
}
+32
View File
@@ -0,0 +1,32 @@
package network
import (
"context"
"errors"
"net"
"time"
)
const (
DefaultTimeout = 10 * time.Second
DefaultIdleTimeout = time.Minute
DefaultHTTPTimeout = 10 * time.Second
DefaultBufferSize = 4096
ProxyDialerOpenThreshold = 5
ProxyDialerHalfOpenTimeout = time.Minute
ProxyDialerResetFailuresTimeout = 10 * time.Second
DefaultDOHHostname = "9.9.9.9"
DNSTimeout = 5 * time.Second
)
var (
ErrCircuitBreakerOpened = errors.New("circuit breaker is opened")
ErrCannotDialWithAllProxies = errors.New("cannot dial with all proxies")
)
type Dialer interface {
Dial(network, address string) (net.Conn, error)
DialContext(ctx context.Context, network, address string) (net.Conn, error)
}
+24
View File
@@ -0,0 +1,24 @@
package network
import (
"context"
"net"
"github.com/stretchr/testify/mock"
)
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)
}
+87
View File
@@ -0,0 +1,87 @@
package network_test
import (
"context"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"github.com/9seconds/mtg/v2/network"
socks5 "github.com/armon/go-socks5"
"github.com/mccutchen/go-httpbin/httpbin"
"github.com/stretchr/testify/mock"
)
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
}
func (suite *HTTPServerTestSuite) SetupSuite() {
suite.httpServer = httptest.NewServer(httpbin.NewHTTPBin().Handler())
}
func (suite *HTTPServerTestSuite) TearDownSuite() {
suite.httpServer.Close()
}
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) // nolint: errcheck
}
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(),
}
}
+50
View File
@@ -0,0 +1,50 @@
package network
import (
"context"
"fmt"
"math/rand"
"net"
"net/url"
)
type loadBalancedSocks5Dialer struct {
dialers []Dialer
}
func (l loadBalancedSocks5Dialer) Dial(network, address string) (net.Conn, error) {
return l.DialContext(context.Background(), network, address)
}
func (l loadBalancedSocks5Dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
length := len(l.dialers)
start := rand.Intn(length)
moved := false
for i := start; i != start || !moved; i = (i + 1) % length {
moved = true
if conn, err := l.dialers[i].DialContext(ctx, network, address); err == nil {
return conn, nil
}
}
return nil, ErrCannotDialWithAllProxies
}
func NewLoadBalancedSocks5Dialer(baseDialer Dialer, proxyURLs []*url.URL) (Dialer, error) {
dialers := make([]Dialer, 0, len(proxyURLs))
for _, u := range proxyURLs {
dialer, err := NewSocks5Dialer(newProxyDialer(baseDialer, u), u)
if err != nil {
return nil, fmt.Errorf("cannot build dialer for %s: %w", u.String(), err)
}
dialers = append(dialers, dialer)
}
return loadBalancedSocks5Dialer{
dialers: dialers,
}, nil
}
+88
View File
@@ -0,0 +1,88 @@
package network_test
import (
"errors"
"io"
"net"
"net/http"
"net/url"
"testing"
"github.com/9seconds/mtg/v2/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")) // nolint: noctx
if err == nil {
defer resp.Body.Close()
}
suite.NoError(err)
suite.Equal(http.StatusOK, resp.StatusCode)
}
func TestLoadBalancedSocks5(t *testing.T) {
t.Parallel()
suite.Run(t, &LoadBalancedSocks5TestSuite{})
}
+171
View File
@@ -0,0 +1,171 @@
package network
import (
"context"
"fmt"
"math/rand"
"net"
"net/http"
"sync"
"time"
"github.com/9seconds/mtg/v2/mtglib"
doh "github.com/babolivier/go-doh-client"
)
type networkHTTPTransport struct {
userAgent string
next http.RoundTripper
}
func (n networkHTTPTransport) RoundTrip(req *http.Request) (*http.Response, error) {
req.Header.Set("User-Agent", n.userAgent)
return n.next.RoundTrip(req)
}
type network struct {
dialer Dialer
dns doh.Resolver
idleTimeout time.Duration
httpTimeout time.Duration
userAgent string
}
func (n *network) Dial(protocol, address string) (net.Conn, error) {
return n.DialContext(context.Background(), protocol, address)
}
func (n *network) DialContext(ctx context.Context, protocol, address string) (net.Conn, error) {
host, port, _ := net.SplitHostPort(address)
ips, err := n.dnsResolve(protocol, host)
if err != nil {
return nil, fmt.Errorf("cannot resolve dns names: %w", err)
}
if len(ips) > 1 {
rand.Shuffle(len(ips), func(i, j int) {
ips[i], ips[j] = ips[j], ips[i]
})
}
var conn net.Conn
for _, v := range ips {
conn, err = n.dialer.DialContext(ctx, protocol, net.JoinHostPort(v, port))
if err == nil {
return conn, nil
}
}
return nil, fmt.Errorf("cannot dial to %s:%s: %w", protocol, address, err)
}
func (n *network) MakeHTTPClient(dialFunc func(ctx context.Context,
network, address string) (net.Conn, error)) *http.Client {
if dialFunc == nil {
dialFunc = n.DialContext
}
return makeHTTPClient(n.userAgent, n.httpTimeout, dialFunc)
}
func (n *network) IdleTimeout() time.Duration {
return n.idleTimeout
}
func (n *network) dnsResolve(protocol, address string) ([]string, error) {
if net.ParseIP(address) != nil {
return []string{address}, nil
}
ips := []string{}
wg := &sync.WaitGroup{}
mutex := &sync.Mutex{}
switch protocol {
case "tcp", "tcp4":
wg.Add(1)
go func() {
defer wg.Done()
if recs, _, err := n.dns.LookupA(address); err == nil {
mutex.Lock()
defer mutex.Unlock()
for _, v := range recs {
ips = append(ips, v.IP4)
}
}
}()
}
switch protocol {
case "tcp", "tcp6":
wg.Add(1)
go func() {
defer wg.Done()
if recs, _, err := n.dns.LookupAAAA(address); err == nil {
mutex.Lock()
defer mutex.Unlock()
for _, v := range recs {
ips = append(ips, v.IP6)
}
}
}()
}
wg.Wait()
if len(ips) == 0 {
return nil, fmt.Errorf("cannot find any ips for %s:%s", protocol, address)
}
return ips, nil
}
func NewNetwork(dialer Dialer,
userAgent, dohHostname string,
httpTimeout, idleTimeout time.Duration) (mtglib.Network, error) {
switch {
case idleTimeout < 0:
return nil, fmt.Errorf("timeout should be positive number %s", idleTimeout)
case idleTimeout == 0:
idleTimeout = DefaultIdleTimeout
}
if net.ParseIP(dohHostname) == nil {
return nil, fmt.Errorf("hostname %s should be IP address", dohHostname)
}
return &network{
dialer: dialer,
idleTimeout: idleTimeout,
httpTimeout: httpTimeout,
userAgent: userAgent,
dns: doh.Resolver{
Host: dohHostname,
Class: doh.IN,
HTTPClient: makeHTTPClient(userAgent, DNSTimeout, dialer.DialContext),
},
}, nil
}
func makeHTTPClient(userAgent string,
timeout time.Duration,
dialFunc func(ctx context.Context, network, address string) (net.Conn, error)) *http.Client {
return &http.Client{
Timeout: timeout,
Transport: networkHTTPTransport{
userAgent: userAgent,
next: &http.Transport{
DialContext: dialFunc,
},
},
}
}
+37
View File
@@ -0,0 +1,37 @@
package network
import (
"net/url"
"strconv"
"time"
)
func newProxyDialer(baseDialer Dialer, proxyURL *url.URL) Dialer {
params := proxyURL.Query()
var (
openThreshold uint32 = ProxyDialerOpenThreshold
halfOpenTimeout = ProxyDialerHalfOpenTimeout
resetFailuresTimeout = ProxyDialerResetFailuresTimeout
)
if param := params.Get("open_threshold"); param != "" {
if intNum, err := strconv.ParseUint(param, 10, 32); err == nil {
openThreshold = uint32(intNum)
}
}
if param := params.Get("half_open_timeout"); param != "" {
if dur, err := time.ParseDuration(param); err == nil && dur > 0 {
halfOpenTimeout = dur
}
}
if param := params.Get("reset_failures_timeout"); param != "" {
if dur, err := time.ParseDuration(param); err == nil && dur > 0 {
resetFailuresTimeout = dur
}
}
return newCircuitBreakerDialer(baseDialer, openThreshold, halfOpenTimeout, resetFailuresTimeout)
}
+94
View File
@@ -0,0 +1,94 @@
package network
import (
"net/url"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/suite"
)
type ProxyDialerTestSuite struct {
suite.Suite
u *url.URL
}
func (suite *ProxyDialerTestSuite) SetupSuite() {
u, _ := url.Parse("socks5://hello:world@10.0.0.10:3128")
suite.u = u
}
func (suite *ProxyDialerTestSuite) TestSetupDefaults() {
d := newProxyDialer(&DialerMock{}, suite.u).(*circuitBreakerDialer)
suite.EqualValues(ProxyDialerOpenThreshold, d.openThreshold)
suite.EqualValues(ProxyDialerHalfOpenTimeout, d.halfOpenTimeout)
suite.EqualValues(ProxyDialerResetFailuresTimeout, d.resetFailuresTimeout)
}
func (suite *ProxyDialerTestSuite) TestSetupValuesAllOk() {
query := url.Values{}
query.Set("open_threshold", "30")
query.Set("reset_failures_timeout", "1s")
query.Set("half_open_timeout", "2s")
suite.u.RawQuery = query.Encode()
d := newProxyDialer(&DialerMock{}, suite.u).(*circuitBreakerDialer)
suite.EqualValues(30, d.openThreshold)
suite.EqualValues(2*time.Second, d.halfOpenTimeout)
suite.EqualValues(time.Second, d.resetFailuresTimeout)
}
func (suite *ProxyDialerTestSuite) TestOpenThreshold() {
query := url.Values{}
params := []string{"-30", "aaa", "1.0", "-1.0"}
for _, v := range params {
param := v
suite.T().Run(v, func(t *testing.T) {
query.Set("open_threshold", param)
suite.u.RawQuery = query.Encode()
d := newProxyDialer(&DialerMock{}, suite.u).(*circuitBreakerDialer)
assert.EqualValues(t, ProxyDialerOpenThreshold, d.openThreshold)
})
}
}
func (suite *ProxyDialerTestSuite) TestHalfOpenTimeout() {
query := url.Values{}
params := []string{"-30", "30", "aaa", "-3.0", "3.0"}
for _, v := range params {
param := v
suite.T().Run(v, func(t *testing.T) {
query.Set("half_open_timeout", param)
suite.u.RawQuery = query.Encode()
d := newProxyDialer(&DialerMock{}, suite.u).(*circuitBreakerDialer)
assert.EqualValues(t, ProxyDialerHalfOpenTimeout, d.halfOpenTimeout)
})
}
}
func (suite *ProxyDialerTestSuite) TestResetFailuresTimeout() {
query := url.Values{}
params := []string{"-30", "30", "aaa", "-3.0", "3.0"}
for _, v := range params {
param := v
suite.T().Run(v, func(t *testing.T) {
query.Set("reset_failures_timeout", param)
suite.u.RawQuery = query.Encode()
d := newProxyDialer(&DialerMock{}, suite.u).(*circuitBreakerDialer)
assert.EqualValues(t, ProxyDialerHalfOpenTimeout, d.halfOpenTimeout)
})
}
}
func TestProxyDialer(t *testing.T) {
t.Parallel()
suite.Run(t, &ProxyDialerTestSuite{})
}
+17
View File
@@ -0,0 +1,17 @@
package network
import (
"fmt"
"net/url"
"golang.org/x/net/proxy"
)
func NewSocks5Dialer(baseDialer Dialer, proxyURL *url.URL) (Dialer, error) {
rv, err := proxy.FromURL(proxyURL, baseDialer)
if err != nil {
return nil, fmt.Errorf("cannot initialize socks5 proxy dialer: %w", err)
}
return rv.(Dialer), nil
}
+61
View File
@@ -0,0 +1,61 @@
package network_test
import (
"net/http"
"testing"
"github.com/9seconds/mtg/v2/network"
"github.com/stretchr/testify/suite"
)
type Socks5TestSuite struct {
suite.Suite
HTTPServerTestSuite
Socks5ServerTestSuite
d network.Dialer
}
func (suite *Socks5TestSuite) SetupSuite() {
suite.HTTPServerTestSuite.SetupSuite()
suite.Socks5ServerTestSuite.SetupSuite()
suite.d, _ = network.NewDefaultDialer(0, 0)
}
func (suite *Socks5TestSuite) TearDownSuite() {
suite.Socks5ServerTestSuite.TearDownSuite()
suite.HTTPServerTestSuite.TearDownSuite()
}
func (suite *Socks5TestSuite) TestRequestFailed() {
proxyURL := suite.MakeSocks5URL("user2", "password")
dialer, _ := network.NewSocks5Dialer(suite.d, proxyURL)
httpClient := suite.MakeHTTPClient(dialer)
resp, err := httpClient.Get(suite.MakeURL("/get")) // nolint: noctx
if err == nil {
defer resp.Body.Close()
}
suite.Error(err)
}
func (suite *Socks5TestSuite) TestRequestOk() {
proxyURL := suite.MakeSocks5URL("user", "password")
dialer, _ := network.NewSocks5Dialer(suite.d, proxyURL)
httpClient := suite.MakeHTTPClient(dialer)
resp, err := httpClient.Get(suite.MakeURL("/get")) // nolint: noctx
if err == nil {
defer resp.Body.Close()
}
suite.NoError(err)
suite.Equal(http.StatusOK, resp.StatusCode)
}
func TestSocks5TestSuite(t *testing.T) {
t.Parallel()
suite.Run(t, &Socks5TestSuite{})
}