From 7002e4cd09b782de66877568b8aa2eeeff4625f1 Mon Sep 17 00:00:00 2001 From: 9seconds Date: Tue, 9 Mar 2021 11:52:13 +0300 Subject: [PATCH] Add load balancing network dialer --- mtglib/network/circuit_breaker.go | 195 ++++++++++++++++++++++++++++++ mtglib/network/consts.go | 9 -- mtglib/network/init.go | 25 ++++ mtglib/network/init_test.go | 19 +++ mtglib/network/interfaces.go | 11 -- mtglib/network/load_balanced.go | 55 +++++++++ mtglib/network/network.go | 2 +- mtglib/network/proxy_dialer.go | 43 +++++++ mtglib/network/socks5.go | 4 +- mtglib/network/socks5_test.go | 4 +- 10 files changed, 342 insertions(+), 25 deletions(-) create mode 100644 mtglib/network/circuit_breaker.go delete mode 100644 mtglib/network/consts.go create mode 100644 mtglib/network/init.go delete mode 100644 mtglib/network/interfaces.go create mode 100644 mtglib/network/load_balanced.go create mode 100644 mtglib/network/proxy_dialer.go diff --git a/mtglib/network/circuit_breaker.go b/mtglib/network/circuit_breaker.go new file mode 100644 index 0000000..4581cb7 --- /dev/null +++ b/mtglib/network/circuit_breaker.go @@ -0,0 +1,195 @@ +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 + } + + c.failuresCount++ + + if c.state == circuitBreakerStateClosed && c.failuresCount > c.openThreshold { + c.switchState(circuitBreakerStateOpened) + } + + return conn, err +} + +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 + } + + if err == nil { + c.switchState(circuitBreakerStateClosed) + } else { + c.switchState(circuitBreakerStateOpened) + } + + return conn, err +} + +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, + openThreshold: openThreshold, + halfOpenTimeout: halfOpenTimeout, + resetFailuresTimeout: resetFailuresTimeout, + } + + cb.switchState(circuitBreakerStateClosed) + + return cb +} diff --git a/mtglib/network/consts.go b/mtglib/network/consts.go deleted file mode 100644 index 68f00fc..0000000 --- a/mtglib/network/consts.go +++ /dev/null @@ -1,9 +0,0 @@ -package network - -import "time" - -const ( - DefaultTimeout = 10 * time.Second - DefaultHTTPTimeout = DefaultTimeout - DefaultBufferSize = 4096 -) diff --git a/mtglib/network/init.go b/mtglib/network/init.go new file mode 100644 index 0000000..9d6a9e4 --- /dev/null +++ b/mtglib/network/init.go @@ -0,0 +1,25 @@ +package network + +import ( + "context" + "errors" + "net" + "time" +) + +const ( + DefaultTimeout = 10 * time.Second + DefaultDNSTimeout = time.Second + DefaultHTTPTimeout = DefaultTimeout + DefaultBufferSize = 4096 +) + +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) +} diff --git a/mtglib/network/init_test.go b/mtglib/network/init_test.go index 8f8aeff..89aba2b 100644 --- a/mtglib/network/init_test.go +++ b/mtglib/network/init_test.go @@ -1,13 +1,32 @@ package network_test import ( + "context" + "net" "net/http/httptest" "strings" "github.com/mccutchen/go-httpbin/httpbin" + "github.com/stretchr/testify/mock" "github.com/stretchr/testify/suite" ) +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 { suite.Suite diff --git a/mtglib/network/interfaces.go b/mtglib/network/interfaces.go deleted file mode 100644 index 2f450ae..0000000 --- a/mtglib/network/interfaces.go +++ /dev/null @@ -1,11 +0,0 @@ -package network - -import ( - "context" - "net" -) - -type Dialer interface { - Dial(network, address string) (net.Conn, error) - DialContext(ctx context.Context, network, address string) (net.Conn, error) -} diff --git a/mtglib/network/load_balanced.go b/mtglib/network/load_balanced.go new file mode 100644 index 0000000..9ffbd82 --- /dev/null +++ b/mtglib/network/load_balanced.go @@ -0,0 +1,55 @@ +package network + +import ( + "context" + "math/rand" + "net" + "net/url" +) + +type loadBalancedDialer struct { + dialers []Dialer +} + +func (l loadBalancedDialer) Dial(network, address string) (net.Conn, error) { + return l.DialContext(context.Background(), network, address) +} + +func (l loadBalancedDialer) 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 NewLoadBalancedDialer(baseDialer Dialer, proxyURLs []*url.URL) (Dialer, error) { + switch len(proxyURLs) { + case 0: + return baseDialer, nil + case 1: + return NewSocks5Dialer(baseDialer, proxyURLs[0]) + } + + dialers := []Dialer{} + + for _, u := range proxyURLs { + dialer, err := NewSocks5Dialer(newProxyDialer(baseDialer, u), u) + if err != nil { + return nil, err + } + + dialers = append(dialers, dialer) + } + + return loadBalancedDialer{ + dialers: dialers, + }, nil +} diff --git a/mtglib/network/network.go b/mtglib/network/network.go index ecbf0e7..31289ff 100644 --- a/mtglib/network/network.go +++ b/mtglib/network/network.go @@ -113,7 +113,7 @@ func NewNetwork(dialer Dialer, dohHostname string, httpTimeout time.Duration) (* } dohHTTPClient := &http.Client{ - Timeout: httpTimeout, + Timeout: DefaultDNSTimeout, Transport: &http.Transport{ DialContext: dialer.DialContext, }, diff --git a/mtglib/network/proxy_dialer.go b/mtglib/network/proxy_dialer.go new file mode 100644 index 0000000..0087819 --- /dev/null +++ b/mtglib/network/proxy_dialer.go @@ -0,0 +1,43 @@ +package network + +import ( + "net/url" + "strconv" + "time" +) + +const ( + ProxyDialerOpenThreshold = 5 + ProxyDialerHalfOpenTimeout = time.Minute + ProxyDialerResetFailuresTimeout = 10 * time.Second +) + +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) +} diff --git a/mtglib/network/socks5.go b/mtglib/network/socks5.go index af4df25..91c9c61 100644 --- a/mtglib/network/socks5.go +++ b/mtglib/network/socks5.go @@ -7,8 +7,8 @@ import ( "golang.org/x/net/proxy" ) -func NewSocks5Dialer(proxyURL *url.URL, base Dialer) (Dialer, error) { - rv, err := proxy.FromURL(proxyURL, base) +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) } diff --git a/mtglib/network/socks5_test.go b/mtglib/network/socks5_test.go index 0b1ec34..f419a69 100644 --- a/mtglib/network/socks5_test.go +++ b/mtglib/network/socks5_test.go @@ -47,7 +47,7 @@ func (suite *Socks5TestSuite) TestRequestFailed() { User: url.UserPassword("user2", "password"), Host: suite.socksListener.Addr().String(), } - dialer, _ := network.NewSocks5Dialer(proxyURL, suite.baseDialer) + dialer, _ := network.NewSocks5Dialer(suite.baseDialer, proxyURL) httpClient := http.Client{ Transport: &http.Transport{ @@ -66,7 +66,7 @@ func (suite *Socks5TestSuite) TestRequestOk() { User: url.UserPassword("user", "password"), Host: suite.socksListener.Addr().String(), } - dialer, _ := network.NewSocks5Dialer(proxyURL, suite.baseDialer) + dialer, _ := network.NewSocks5Dialer(suite.baseDialer, proxyURL) httpClient := http.Client{ Transport: &http.Transport{