mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 14:54:01 +03:00
Refactor network to a top-level module
This commit is contained in:
+28
-1
@@ -1,5 +1,32 @@
|
||||
package mtglib
|
||||
|
||||
import "errors"
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
var ErrSecretEmpty = errors.New("secret is empty")
|
||||
|
||||
type Network interface {
|
||||
Dial(network, address string) (net.Conn, error)
|
||||
DialContext(ctx context.Context, network, address string) (net.Conn, error)
|
||||
MakeHTTPClient(func(ctx context.Context, network, address string) (net.Conn, error)) *http.Client
|
||||
IdleTimeout() time.Duration
|
||||
}
|
||||
|
||||
type Logger interface {
|
||||
Named(name string) Logger
|
||||
|
||||
BindInt(name string, value int) Logger
|
||||
BindStr(name, value string) Logger
|
||||
|
||||
Info(msg string)
|
||||
InfoError(msg string, err error)
|
||||
Warning(msg string)
|
||||
WarningError(msg string, err error)
|
||||
Debug(msg string)
|
||||
DebugError(msg string, err error)
|
||||
}
|
||||
|
||||
@@ -1,196 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,138 +0,0 @@
|
||||
package network
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"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 *ConnMock
|
||||
baseDialerMock *DialerMock
|
||||
}
|
||||
|
||||
func (suite *CircuitBreakerTestSuite) SetupTest() {
|
||||
suite.mutex = sync.Mutex{}
|
||||
suite.ctx, suite.ctxCancel = context.WithCancel(context.Background())
|
||||
suite.baseDialerMock = &DialerMock{}
|
||||
suite.connMock = &ConnMock{}
|
||||
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{})
|
||||
}
|
||||
@@ -1,86 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,77 +0,0 @@
|
||||
package network_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/9seconds/mtg/v2/mtglib/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{})
|
||||
}
|
||||
@@ -1,44 +0,0 @@
|
||||
package network
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"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 DialFunc func(ctx context.Context, protocol, address string) (net.Conn, error)
|
||||
|
||||
type Dialer interface {
|
||||
Dial(network, address string) (net.Conn, error)
|
||||
DialContext(ctx context.Context, network, address string) (net.Conn, error)
|
||||
}
|
||||
|
||||
type Network interface {
|
||||
Dialer
|
||||
|
||||
DNSResolve(network, hostname string) (ips []string, err error)
|
||||
MakeHTTPClient(DialFunc) *http.Client
|
||||
IdleTimeout() time.Duration
|
||||
HTTPTimeout() time.Duration
|
||||
}
|
||||
@@ -1,65 +0,0 @@
|
||||
package network
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/mock"
|
||||
)
|
||||
|
||||
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)
|
||||
}
|
||||
@@ -1,87 +0,0 @@
|
||||
package network_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/9seconds/mtg/v2/mtglib/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(),
|
||||
}
|
||||
}
|
||||
@@ -1,50 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,88 +0,0 @@
|
||||
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")) // 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{})
|
||||
}
|
||||
@@ -1,171 +0,0 @@
|
||||
package network
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
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) 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 (n *network) MakeHTTPClient(dialFunc DialFunc) *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) HTTPTimeout() time.Duration {
|
||||
return n.httpTimeout
|
||||
}
|
||||
|
||||
func NewNetwork(dialer Dialer,
|
||||
userAgent, dohHostname string,
|
||||
httpTimeout, idleTimeout time.Duration) (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 DialFunc) *http.Client {
|
||||
return &http.Client{
|
||||
Timeout: timeout,
|
||||
Transport: networkHTTPTransport{
|
||||
userAgent: userAgent,
|
||||
next: &http.Transport{
|
||||
DialContext: dialFunc,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -1,37 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -1,94 +0,0 @@
|
||||
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{})
|
||||
}
|
||||
@@ -1,17 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,61 +0,0 @@
|
||||
package network_test
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/9seconds/mtg/v2/mtglib/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{})
|
||||
}
|
||||
Reference in New Issue
Block a user