Add load balancing network dialer

This commit is contained in:
9seconds
2021-03-09 11:52:13 +03:00
parent ecfd550a95
commit 7002e4cd09
10 changed files with 342 additions and 25 deletions
+195
View File
@@ -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
}
-9
View File
@@ -1,9 +0,0 @@
package network
import "time"
const (
DefaultTimeout = 10 * time.Second
DefaultHTTPTimeout = DefaultTimeout
DefaultBufferSize = 4096
)
+25
View File
@@ -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)
}
+19
View File
@@ -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
-11
View File
@@ -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)
}
+55
View File
@@ -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
}
+1 -1
View File
@@ -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,
},
+43
View File
@@ -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)
}
+2 -2
View File
@@ -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)
}
+2 -2
View File
@@ -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{