Change algorithm of TCP relaying

This commit is contained in:
9seconds
2021-08-27 17:22:36 +03:00
parent 4b7be8c565
commit 456ed5b051
18 changed files with 300 additions and 434 deletions
+1 -2
View File
@@ -183,7 +183,6 @@ func runProxy(conf *config.Config, version string) error {
Secret: conf.Secret, Secret: conf.Secret,
BufferSize: conf.TCPBuffer.Get(mtglib.DefaultBufferSize), BufferSize: conf.TCPBuffer.Get(mtglib.DefaultBufferSize),
DomainFrontingPort: conf.DomainFrontingPort.Get(mtglib.DefaultDomainFrontingPort), DomainFrontingPort: conf.DomainFrontingPort.Get(mtglib.DefaultDomainFrontingPort),
IdleTimeout: conf.Network.Timeout.Idle.Get(mtglib.DefaultIdleTimeout),
PreferIP: conf.PreferIP.Get(mtglib.DefaultPreferIP), PreferIP: conf.PreferIP.Get(mtglib.DefaultPreferIP),
} }
@@ -192,7 +191,7 @@ func runProxy(conf *config.Config, version string) error {
return fmt.Errorf("cannot create a proxy: %w", err) return fmt.Errorf("cannot create a proxy: %w", err)
} }
listener, err := net.Listen("tcp", conf.BindTo.Get("")) listener, err := utils.NewListener(conf.BindTo.Get(""), int(opts.BufferSize))
if err != nil { if err != nil {
return fmt.Errorf("cannot start proxy: %w", err) return fmt.Errorf("cannot start proxy: %w", err)
} }
+41
View File
@@ -0,0 +1,41 @@
package utils
import (
"fmt"
"net"
"github.com/9seconds/mtg/v2/network"
)
type Listener struct {
net.Listener
bufferSize int
}
func (l Listener) Accept() (net.Conn, error) {
conn, err := l.Listener.Accept()
if err != nil {
return nil, err // nolint: wrapcheck
}
if err := network.SetClientSocketOptions(conn, l.bufferSize); err != nil {
conn.Close()
return nil, fmt.Errorf("cannot set TCP options: %w", err)
}
return conn, nil
}
func NewListener(bindTo string, bufferSize int) (net.Listener, error) {
base, err := net.Listen("tcp", bindTo)
if err != nil {
return nil, fmt.Errorf("cannot build a base listener: %w", err)
}
return Listener{
Listener: base,
bufferSize: bufferSize,
}, nil
}
+10 -2
View File
@@ -69,6 +69,9 @@ const (
// DefaultIdleTimeout is a default timeout for closing a connection // DefaultIdleTimeout is a default timeout for closing a connection
// in case of idling. // in case of idling.
//
// Deprecated: no longer in use because of changed TCP relay
// algorithm.
DefaultIdleTimeout = time.Minute DefaultIdleTimeout = time.Minute
// DefaultTolerateTimeSkewness is a default timeout for time // DefaultTolerateTimeSkewness is a default timeout for time
@@ -83,9 +86,14 @@ const (
// by Telegram and a proxy. // by Telegram and a proxy.
SecretKeyLength = 16 SecretKeyLength = 16
// ConnectionIDBytesLength defines a count of random bytes // ConnectionIDBytesLength defines a count of random bytes used to generate
// used to generate a stream/connection ids. // a stream/connection ids.
ConnectionIDBytesLength = 16 ConnectionIDBytesLength = 16
// TCPRelayReadTimeout defines a max time period between two consecuitive
// reads from Telegram after which connection will be terminated. This is
// required to abort stale connections.
TCPRelayReadTimeout = 20 * time.Second
) )
// Network defines a knowledge how to work with a network. It may sound // Network defines a knowledge how to work with a network. It may sound
+11 -7
View File
@@ -46,7 +46,11 @@ func (c *Conn) Write(p []byte) (int, error) {
rec.Type = record.TypeApplicationData rec.Type = record.TypeApplicationData
rec.Version = record.Version12 rec.Version = record.Version12
written := 0
sendBuffer := acquireBytesBuffer()
defer releaseBytesBuffer(sendBuffer)
lenP := len(p)
for len(p) > 0 { for len(p) > 0 {
chunkSize := rand.Intn(record.TLSMaxRecordSize) chunkSize := rand.Intn(record.TLSMaxRecordSize)
@@ -56,14 +60,14 @@ func (c *Conn) Write(p []byte) (int, error) {
rec.Payload.Reset() rec.Payload.Reset()
rec.Payload.Write(p[:chunkSize]) rec.Payload.Write(p[:chunkSize])
rec.Dump(sendBuffer) // nolint: errcheck
if err := rec.Dump(c.Conn); err != nil {
return written, err // nolint: wrapcheck
}
written += chunkSize
p = p[chunkSize:] p = p[chunkSize:]
} }
return written, nil if _, err := c.Conn.Write(sendBuffer.Bytes()); err != nil {
return 0, err // nolint: wrapcheck
}
return lenP, nil
} }
+7 -23
View File
@@ -1,35 +1,19 @@
package relay package relay
import ( import (
"context" "fmt"
"io" "net"
"time"
) )
type conn struct { type conn struct {
io.ReadWriteCloser net.Conn
ctx context.Context
tickChannel chan struct{}
} }
func (c conn) Read(p []byte) (int, error) { func (c conn) Read(p []byte) (int, error) {
n, err := c.ReadWriteCloser.Read(p) if err := c.SetReadDeadline(time.Now().Add(getTimeout())); err != nil {
return 0, fmt.Errorf("cannot set read deadline: %w", err)
select {
case <-c.ctx.Done():
case c.tickChannel <- struct{}{}:
} }
return n, err // nolint: wrapcheck return c.Conn.Read(p) // nolint: wrapcheck
}
func (c conn) Write(p []byte) (int, error) {
n, err := c.ReadWriteCloser.Write(p)
select {
case <-c.ctx.Done():
case c.tickChannel <- struct{}{}:
}
return n, err // nolint: wrapcheck
} }
-125
View File
@@ -1,125 +0,0 @@
package relay
import (
"context"
"errors"
"io"
"testing"
"github.com/9seconds/mtg/v2/internal/testlib"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/suite"
)
type ConnTestSuite struct {
suite.Suite
ctxCancel context.CancelFunc
connMock *testlib.NetConnMock
tickChannel chan struct{}
buf []byte
c conn
}
func (suite *ConnTestSuite) SetupTest() {
ctx, cancel := context.WithCancel(context.Background())
suite.tickChannel = make(chan struct{}, 1)
suite.connMock = &testlib.NetConnMock{}
suite.ctxCancel = cancel
suite.buf = make([]byte, 5)
suite.c = conn{
ReadWriteCloser: suite.connMock,
ctx: ctx,
tickChannel: suite.tickChannel,
}
}
func (suite *ConnTestSuite) TestReadOk() {
suite.connMock.On("Read", mock.Anything).Once().Return(len(suite.buf), nil)
n, err := suite.c.Read(suite.buf)
suite.NoError(err)
suite.Equal(len(suite.buf), n)
select {
case <-suite.tickChannel:
default:
suite.FailNow("cannot find a tick event")
}
}
func (suite *ConnTestSuite) TestReadErr() {
suite.connMock.On("Read", mock.Anything).Once().Return(0, io.EOF)
_, err := suite.c.Read(suite.buf)
suite.True(errors.Is(err, io.EOF))
select {
case <-suite.tickChannel:
default:
suite.FailNow("cannot find a tick event")
}
}
func (suite *ConnTestSuite) TestReadContextDone() {
suite.connMock.On("Read", mock.Anything).Once().Return(len(suite.buf), nil)
suite.ctxCancel()
suite.tickChannel <- struct{}{}
suite.c.Read(suite.buf) // nolint: errcheck
}
func (suite *ConnTestSuite) TestWriteOk() {
suite.connMock.On("Write", mock.Anything).Once().Return(len(suite.buf), nil)
n, err := suite.c.Write(suite.buf)
suite.NoError(err)
suite.Equal(len(suite.buf), n)
select {
case <-suite.tickChannel:
default:
suite.FailNow("cannot find a tick event")
}
}
func (suite *ConnTestSuite) TestWriteErr() {
suite.connMock.On("Write", mock.Anything).Once().Return(0, io.EOF)
_, err := suite.c.Write(suite.buf)
suite.True(errors.Is(err, io.EOF))
select {
case <-suite.tickChannel:
default:
suite.FailNow("cannot find a tick event")
}
}
func (suite *ConnTestSuite) TestWriteContextDone() {
suite.connMock.On("Write", mock.Anything).Once().Return(len(suite.buf), nil)
suite.ctxCancel()
suite.tickChannel <- struct{}{}
suite.c.Write(suite.buf) // nolint: errcheck
}
func (suite *ConnTestSuite) TearDownTest() {
select {
case <-suite.tickChannel:
default:
}
close(suite.tickChannel)
suite.connMock.AssertExpectations(suite.T())
}
func TestConn(t *testing.T) {
t.Parallel()
suite.Run(t, &ConnTestSuite{})
}
+9
View File
@@ -1,5 +1,14 @@
package relay package relay
import "time"
const (
ConnectionTimeToLiveMin = 2 * time.Minute
ConnectionTimeToLiveMax = 10 * time.Minute
TimeoutMin = 20 * time.Second
TimeoutMax = time.Minute
)
type Logger interface { type Logger interface {
Printf(msg string, args ...interface{}) Printf(msg string, args ...interface{})
} }
-44
View File
@@ -1,49 +1,5 @@
package relay_test package relay_test
import (
"bytes"
"io"
"sync"
)
type loggerMock struct{} type loggerMock struct{}
func (l loggerMock) Printf(format string, args ...interface{}) {} func (l loggerMock) Printf(format string, args ...interface{}) {}
type rwcMock struct {
bytes.Buffer
closed bool
mutex sync.Mutex
}
func (r *rwcMock) Read(p []byte) (int, error) {
r.mutex.Lock()
defer r.mutex.Unlock()
if r.closed {
return 0, io.EOF
}
return r.Buffer.Read(p) // nolint: wrapcheck
}
func (r *rwcMock) Write(p []byte) (int, error) {
r.mutex.Lock()
defer r.mutex.Unlock()
if r.closed {
return 0, io.EOF
}
return r.Buffer.Write(p) // nolint: wrapcheck
}
func (r *rwcMock) Close() error {
r.mutex.Lock()
defer r.mutex.Unlock()
r.closed = true
return nil
}
+17 -30
View File
@@ -1,45 +1,32 @@
package relay package relay
import ( import "sync"
"context"
"sync"
"time"
)
var relayPool = sync.Pool{ type eastWest struct {
east []byte
west []byte
}
var eastWestPool = sync.Pool{
New: func() interface{} { New: func() interface{} {
return &Relay{ return &eastWest{}
tickChannel: make(chan struct{}),
errorChannel: make(chan error, 1),
}
}, },
} }
func AcquireRelay(ctx context.Context, logger Logger, bufferSize int, idleTimeout time.Duration) *Relay { func acquireEastWest(bufferSize int) *eastWest {
ctx, cancel := context.WithCancel(ctx) wanted := eastWestPool.Get().(*eastWest) // nolint: forcetypeassert
r, ok := relayPool.Get().(*Relay) if len(wanted.east) != bufferSize {
if !ok { wanted.east = make([]byte, bufferSize)
panic("Relay pool has no relay!")
} }
r.ctx = ctx if len(wanted.west) != bufferSize {
r.ctxCancel = cancel wanted.west = make([]byte, bufferSize)
r.logger = logger
r.tickTimeout = idleTimeout
if len(r.eastBuffer) != bufferSize {
r.eastBuffer = make([]byte, bufferSize)
} }
if len(r.westBuffer) != bufferSize { return wanted
r.westBuffer = make([]byte, bufferSize)
}
return r
} }
func ReleaseRelay(r *Relay) { func releaseEastWest(ew *eastWest) {
r.Reset() eastWestPool.Put(ew)
relayPool.Put(r)
} }
+25 -104
View File
@@ -3,127 +3,48 @@ package relay
import ( import (
"context" "context"
"io" "io"
"net"
"sync" "sync"
"time"
) )
type Relay struct { func Relay(ctx context.Context, log Logger, bufferSize int,
ctx context.Context telegramConn net.Conn, clientConn io.ReadWriteCloser) {
ctxCancel context.CancelFunc defer telegramConn.Close()
logger Logger defer clientConn.Close()
processMutex sync.Mutex
eastBuffer []byte
westBuffer []byte
tickChannel chan struct{}
errorChannel chan error
tickTimeout time.Duration
}
func (r *Relay) Reset() { ctx, cancel := context.WithTimeout(ctx, getConnectionTimeToLive())
r.processMutex.Lock() defer cancel()
defer r.processMutex.Unlock()
if r.ctxCancel != nil { go func() {
r.ctxCancel() <-ctx.Done()
} telegramConn.Close()
clientConn.Close()
}()
r.ctx = nil buffers := acquireEastWest(bufferSize)
r.ctxCancel = nil defer releaseEastWest(buffers)
r.logger = nil
}
func (r *Relay) Process(eastConn, westConn io.ReadWriteCloser) error { telegramConn = conn{
r.processMutex.Lock() Conn: telegramConn,
defer r.processMutex.Unlock()
eastConn = conn{
ReadWriteCloser: eastConn,
ctx: r.ctx,
tickChannel: r.tickChannel,
}
westConn = conn{
ReadWriteCloser: westConn,
ctx: r.ctx,
tickChannel: r.tickChannel,
} }
wg := &sync.WaitGroup{} wg := &sync.WaitGroup{}
wg.Add(3) // nolint: gomnd wg.Add(2) // nolint: gomnd
go r.runObserver(eastConn, westConn, wg) go pump(log, telegramConn, clientConn, wg, buffers.east, "east -> west")
go r.transmit(eastConn, westConn, r.westBuffer, "west", wg) pump(log, clientConn, telegramConn, wg, buffers.west, "west -> east")
r.transmit(westConn, eastConn, r.eastBuffer, "east", wg)
wg.Wait() wg.Wait()
select {
case err := <-r.errorChannel:
return err
default:
return nil
}
} }
func (r *Relay) transmit(src io.ReadCloser, dst io.WriteCloser, func pump(log Logger, src io.ReadCloser, dst io.WriteCloser, wg *sync.WaitGroup,
buffer []byte, direction string, wg *sync.WaitGroup) { buf []byte, direction string) {
defer wg.Done() defer wg.Done()
defer src.Close()
defer dst.Close()
defer func() { if n, err := io.CopyBuffer(dst, src, buf); err != nil {
r.ctxCancel() log.Printf("cannot pump %s (written %d bytes): %w", direction, n, err)
src.Close()
dst.Close()
}()
if _, err := io.CopyBuffer(dst, src, buffer); err != nil {
r.logger.Printf("error '%v' happened on direction %s", err, direction)
select {
case <-r.ctx.Done():
err = r.ctx.Err()
default:
}
select {
case r.errorChannel <- err:
default:
}
}
}
func (r *Relay) runObserver(one, another io.Closer, wg *sync.WaitGroup) {
defer wg.Done()
ticker := time.NewTicker(time.Second)
defer func() {
one.Close()
another.Close()
ticker.Stop()
select {
case <-ticker.C:
default:
}
}()
lastTickAt := time.Now()
for {
select {
case <-r.ctx.Done():
return
case <-r.tickChannel:
lastTickAt = time.Now()
case <-ticker.C:
if time.Since(lastTickAt) > r.tickTimeout {
r.logger.Printf("exit due to a timeout")
r.ctxCancel()
return
}
}
} }
} }
+23 -44
View File
@@ -4,7 +4,6 @@ import (
"context" "context"
"io" "io"
"testing" "testing"
"time"
"github.com/9seconds/mtg/v2/internal/testlib" "github.com/9seconds/mtg/v2/internal/testlib"
"github.com/9seconds/mtg/v2/mtglib/internal/relay" "github.com/9seconds/mtg/v2/mtglib/internal/relay"
@@ -15,60 +14,40 @@ import (
type RelayTestSuite struct { type RelayTestSuite struct {
suite.Suite suite.Suite
ctx context.Context loggerMock relay.Logger
ctxCancel context.CancelFunc ctx context.Context
r *relay.Relay ctxCancel context.CancelFunc
telegramConnMock *testlib.NetConnMock
clientConnMock *testlib.NetConnMock
} }
func (suite *RelayTestSuite) SetupTest() { func (suite *RelayTestSuite) SetupTest() {
suite.ctx, suite.ctxCancel = context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
suite.r = relay.AcquireRelay(suite.ctx, loggerMock{}, 4096, time.Second) suite.ctx = ctx
suite.ctxCancel = cancel
suite.loggerMock = &loggerMock{}
suite.telegramConnMock = &testlib.NetConnMock{}
suite.clientConnMock = &testlib.NetConnMock{}
} }
func (suite *RelayTestSuite) TearDownTest() { func (suite *RelayTestSuite) TearDownTest() {
suite.ctxCancel() suite.ctxCancel()
relay.ReleaseRelay(suite.r) suite.telegramConnMock.AssertExpectations(suite.T())
suite.r = nil suite.clientConnMock.AssertExpectations(suite.T())
} }
func (suite *RelayTestSuite) TestCancelled() { func (suite *RelayTestSuite) TestExit() {
suite.ctxCancel() suite.telegramConnMock.On("SetReadDeadline", mock.Anything).Return(nil)
suite.telegramConnMock.On("Close").Return(nil)
suite.telegramConnMock.On("Read", mock.Anything).Return(10, io.EOF).Once()
suite.telegramConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
eastConn := &rwcMock{} suite.clientConnMock.On("Read", mock.Anything).Return(0, io.EOF).Once()
eastConn.Write([]byte{1, 2, 3, 4, 5}) // nolint: errcheck suite.clientConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
suite.clientConnMock.On("Close").Return(nil)
westConn := &rwcMock{} relay.Relay(suite.ctx, suite.loggerMock, 1024,
westConn.Write([]byte{100, 101, 102}) // nolint: errcheck suite.telegramConnMock, suite.clientConnMock)
suite.Nil(suite.r.Process(eastConn, westConn))
}
func (suite *RelayTestSuite) TestCopyFine() {
eastConn := &rwcMock{}
eastConn.Write([]byte{1, 2, 3, 4, 5}) // nolint: errcheck
westConn := &rwcMock{}
westConn.Write([]byte{100, 101, 102}) // nolint: errcheck
// yes, this test is not good enough. but apparently, if it hangs,
// we can debug most of possible issues.
_ = suite.r.Process(eastConn, westConn)
}
func (suite *RelayTestSuite) TestTimeout() {
eastConn := &rwcMock{}
eastConn.Write([]byte{1, 2, 3, 4, 5}) // nolint: errcheck
westConn := &testlib.NetConnMock{}
westConn.On("Close").Return(nil)
westConn.On("Read", mock.Anything).Return(0, io.EOF).Run(func(_ mock.Arguments) {
time.Sleep(2 * time.Second)
})
westConn.On("Write", mock.Anything).Return(0, io.EOF).Run(func(_ mock.Arguments) {
time.Sleep(2 * time.Second)
})
suite.Error(suite.r.Process(eastConn, westConn))
} }
func TestRelay(t *testing.T) { func TestRelay(t *testing.T) {
+26
View File
@@ -0,0 +1,26 @@
package relay
import (
"math"
"math/rand"
"time"
)
func getConnectionTimeToLive() time.Duration {
return getTime(ConnectionTimeToLiveMin, ConnectionTimeToLiveMax)
}
func getTimeout() time.Duration {
return getTime(TimeoutMin, TimeoutMax)
}
func getTime(minDuration, maxDuration time.Duration) time.Duration {
minDurationInSeconds := minDuration.Seconds()
maxDurationInSeconds := maxDuration.Seconds()
middle := minDurationInSeconds + (maxDurationInSeconds-minDurationInSeconds)/2 // nolint: gomnd
number := minDurationInSeconds + rand.ExpFloat64()*middle
number = math.Round(math.Min(maxDurationInSeconds, number))
return time.Duration(number) * time.Second
}
@@ -0,0 +1,37 @@
package relay
import (
"fmt"
"testing"
"github.com/stretchr/testify/suite"
)
type TimeoutsTestSuite struct {
suite.Suite
}
func (suite *TimeoutsTestSuite) TestGetConnectionTimeToLive() {
for i := 0; i < 100; i++ {
value := getConnectionTimeToLive()
message := fmt.Sprintf("generated value is %v", value)
suite.GreaterOrEqual(value, ConnectionTimeToLiveMin, message)
suite.LessOrEqual(value, ConnectionTimeToLiveMax, message)
}
}
func (suite *TimeoutsTestSuite) TestGetTimeout() {
for i := 0; i < 100; i++ {
value := getTimeout()
message := fmt.Sprintf("generated value is %v", value)
suite.GreaterOrEqual(value, TimeoutMin, message)
suite.LessOrEqual(value, TimeoutMax, message)
}
}
func TestTimeouts(t *testing.T) {
t.Parallel()
suite.Run(t, &TimeoutsTestSuite{})
}
+14 -16
View File
@@ -23,7 +23,6 @@ type Proxy struct {
ctxCancel context.CancelFunc ctxCancel context.CancelFunc
streamWaitGroup sync.WaitGroup streamWaitGroup sync.WaitGroup
idleTimeout time.Duration
tolerateTimeSkewness time.Duration tolerateTimeSkewness time.Duration
bufferSize int bufferSize int
domainFrontingPort int domainFrontingPort int
@@ -81,13 +80,13 @@ func (p *Proxy) ServeConn(conn net.Conn) {
return return
} }
rel := relay.AcquireRelay(ctx, relay.Relay(
p.logger.Named("relay"), p.bufferSize, p.idleTimeout) ctx,
defer relay.ReleaseRelay(rel) ctx.logger.Named("relay"),
p.bufferSize,
if err := rel.Process(ctx.clientConn, ctx.telegramConn); err != nil { ctx.telegramConn,
p.logger.DebugError("relay has been finished", err) ctx.clientConn,
} )
} }
// Serve starts a proxy on a given listener. // Serve starts a proxy on a given listener.
@@ -255,13 +254,13 @@ func (p *Proxy) doDomainFronting(ctx *streamContext, conn *connRewind) {
stream: p.eventStream, stream: p.eventStream,
} }
rel := relay.AcquireRelay(ctx, relay.Relay(
p.logger.Named("domain-fronting"), p.bufferSize, p.idleTimeout) ctx,
defer relay.ReleaseRelay(rel) ctx.logger.Named("domain-fronting"),
p.bufferSize,
if err := rel.Process(conn, frontConn); err != nil { frontConn,
p.logger.DebugError("domain fronting relay has been finished", err) conn,
} )
} }
// NewProxy makes a new proxy instance. // NewProxy makes a new proxy instance.
@@ -287,7 +286,6 @@ func NewProxy(opts ProxyOpts) (*Proxy, error) {
logger: opts.getLogger("proxy"), logger: opts.getLogger("proxy"),
domainFrontingPort: opts.getDomainFrontingPort(), domainFrontingPort: opts.getDomainFrontingPort(),
tolerateTimeSkewness: opts.getTolerateTimeSkewness(), tolerateTimeSkewness: opts.getTolerateTimeSkewness(),
idleTimeout: opts.getIdleTimeout(),
bufferSize: opts.getBufferSize(), bufferSize: opts.getBufferSize(),
telegram: tg, telegram: tg,
} }
-8
View File
@@ -143,14 +143,6 @@ func (p ProxyOpts) getDomainFrontingPort() int {
return int(p.DomainFrontingPort) return int(p.DomainFrontingPort)
} }
func (p ProxyOpts) getIdleTimeout() time.Duration {
if p.IdleTimeout == 0 {
return DefaultIdleTimeout
}
return p.IdleTimeout
}
func (p ProxyOpts) getTolerateTimeSkewness() time.Duration { func (p ProxyOpts) getTolerateTimeSkewness() time.Duration {
if p.TolerateTimeSkewness == 0 { if p.TolerateTimeSkewness == 0 {
return DefaultTolerateTimeSkewness return DefaultTolerateTimeSkewness
+4 -29
View File
@@ -5,8 +5,6 @@ import (
"fmt" "fmt"
"net" "net"
"time" "time"
"github.com/libp2p/go-reuseport"
) )
type defaultDialer struct { type defaultDialer struct {
@@ -31,36 +29,14 @@ func (d *defaultDialer) DialContext(ctx context.Context, network, address string
return nil, fmt.Errorf("cannot dial to %s: %w", address, err) return nil, fmt.Errorf("cannot dial to %s: %w", address, err)
} }
tcpConn, ok := conn.(*net.TCPConn) // we do not need to call to end user. End users call us.
if !ok { if err := SetServerSocketOptions(conn, d.bufferSize); err != nil {
panic("conn type is not tcp")
}
if err := tcpConn.SetNoDelay(true); err != nil {
conn.Close() conn.Close()
return nil, fmt.Errorf("cannot set TCP_NO_DELAY: %w", err) return nil, fmt.Errorf("cannot set socket options: %w", err)
} }
if err := tcpConn.SetReadBuffer(d.bufferSize); err != nil { return conn, 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
} }
// NewDefaultDialer build a new dialer which dials bypassing proxies // NewDefaultDialer build a new dialer which dials bypassing proxies
@@ -87,7 +63,6 @@ func NewDefaultDialer(timeout time.Duration, bufferSize int) (Dialer, error) {
return &defaultDialer{ return &defaultDialer{
Dialer: net.Dialer{ Dialer: net.Dialer{
Timeout: timeout, Timeout: timeout,
Control: reuseport.Control,
}, },
bufferSize: bufferSize, bufferSize: bufferSize,
}, nil }, nil
+4
View File
@@ -70,6 +70,10 @@ const (
// DNSTimeout defines a timeout for DNS queries. // DNSTimeout defines a timeout for DNS queries.
DNSTimeout = 5 * time.Second DNSTimeout = 5 * time.Second
// tcpLingerTimeout defines a number of seconds to wait for sending
// unacknowledged data.
tcpLingerTimeout = 1
) )
var ( var (
+71
View File
@@ -0,0 +1,71 @@
package network
import (
"fmt"
"net"
"golang.org/x/sys/unix"
)
// SetClientSocketOptions tunes a TCP socket that represents a connection to
// end user (not Telegram service or fronting domain).
func SetClientSocketOptions(conn net.Conn, bufferSize int) error {
tcpConn := conn.(*net.TCPConn) // nolint: forcetypeassert
if err := tcpConn.SetNoDelay(false); err != nil {
return fmt.Errorf("cannot disable TCP_NO_DELAY: %w", err)
}
return setCommonSocketOptions(tcpConn, bufferSize)
}
// SetServerSocketOptions tunes a TCP socket that represents a connection to
// remote server like Telegram or fronting domain (but not end user).
func SetServerSocketOptions(conn net.Conn, bufferSize int) error {
tcpConn := conn.(*net.TCPConn) // nolint: forcetypeassert
if err := tcpConn.SetNoDelay(true); err != nil {
return fmt.Errorf("cannot enable TCP_NO_DELAY: %w", err)
}
return setCommonSocketOptions(tcpConn, bufferSize)
}
func setCommonSocketOptions(conn *net.TCPConn, bufferSize int) error {
if err := conn.SetReadBuffer(bufferSize); err != nil {
return fmt.Errorf("cannot set read buffer size: %w", err)
}
if err := conn.SetWriteBuffer(bufferSize); err != nil {
return fmt.Errorf("cannot set write buffer size: %w", err)
}
if err := conn.SetKeepAlive(false); err != nil {
return fmt.Errorf("cannot disable TCP keepalive probes: %w", err)
}
if err := conn.SetLinger(tcpLingerTimeout); err != nil {
return fmt.Errorf("cannot set TCP linger timeout: %w", err)
}
rawConn, err := conn.SyscallConn()
if err != nil {
return fmt.Errorf("cannot get underlying raw connection")
}
rawConn.Control(func(fd uintptr) { // nolint: errcheck
err = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEADDR, 1)
if err != nil {
err = fmt.Errorf("cannot set SO_REUSEADDR: %w", err)
return
}
err = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1)
if err != nil {
err = fmt.Errorf("cannot set SO_REUSEPORT: %w", err)
}
})
return nil
}