mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 11:04:02 +03:00
Merge pull request #424 from dolonet/fix/relay-idle-timeout-shared-tracker
fix: use shared idle tracker for relay connections
This commit is contained in:
+53
-5
@@ -6,6 +6,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/9seconds/mtg/v2/essentials"
|
"github.com/9seconds/mtg/v2/essentials"
|
||||||
@@ -97,20 +98,67 @@ func newConnProxyProtocol(source, target essentials.Conn) *connProxyProtocol {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// idleTracker is a shared idle tracker for a pair of relay connections.
|
||||||
|
// Both directions update the same timestamp so that activity in one direction
|
||||||
|
// prevents the other (idle) direction from timing out.
|
||||||
|
type idleTracker struct {
|
||||||
|
lastActive atomic.Int64 // unix nanos
|
||||||
|
timeout time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
func newIdleTracker(timeout time.Duration) *idleTracker {
|
||||||
|
t := &idleTracker{timeout: timeout}
|
||||||
|
t.touch()
|
||||||
|
|
||||||
|
return t
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *idleTracker) touch() {
|
||||||
|
t.lastActive.Store(time.Now().UnixNano())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *idleTracker) isIdle() bool {
|
||||||
|
last := time.Unix(0, t.lastActive.Load())
|
||||||
|
|
||||||
|
return time.Since(last) >= t.timeout
|
||||||
|
}
|
||||||
|
|
||||||
type connIdleTimeout struct {
|
type connIdleTimeout struct {
|
||||||
essentials.Conn
|
essentials.Conn
|
||||||
|
|
||||||
timeout time.Duration
|
tracker *idleTracker
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c connIdleTimeout) Read(b []byte) (int, error) {
|
func (c connIdleTimeout) Read(b []byte) (int, error) {
|
||||||
c.SetReadDeadline(time.Now().Add(c.timeout)) //nolint: errcheck
|
for {
|
||||||
|
c.SetReadDeadline(time.Now().Add(c.tracker.timeout)) //nolint: errcheck
|
||||||
|
|
||||||
return c.Conn.Read(b) //nolint: wrapcheck
|
n, err := c.Conn.Read(b)
|
||||||
|
if n > 0 {
|
||||||
|
c.tracker.touch()
|
||||||
|
|
||||||
|
return n, err //nolint: wrapcheck
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
if netErr, ok := err.(net.Error); ok && netErr.Timeout() && !c.tracker.isIdle() { //nolint: errorlint
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
return 0, err //nolint: wrapcheck
|
||||||
|
}
|
||||||
|
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c connIdleTimeout) Write(b []byte) (int, error) {
|
func (c connIdleTimeout) Write(b []byte) (int, error) {
|
||||||
c.SetWriteDeadline(time.Now().Add(c.timeout)) //nolint: errcheck
|
c.SetWriteDeadline(time.Now().Add(c.tracker.timeout)) //nolint: errcheck
|
||||||
|
|
||||||
return c.Conn.Write(b) //nolint: wrapcheck
|
n, err := c.Conn.Write(b)
|
||||||
|
if n > 0 {
|
||||||
|
c.tracker.touch()
|
||||||
|
}
|
||||||
|
|
||||||
|
return n, err //nolint: wrapcheck
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,6 +16,12 @@ import (
|
|||||||
"github.com/stretchr/testify/suite"
|
"github.com/stretchr/testify/suite"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type netTimeoutError struct{}
|
||||||
|
|
||||||
|
func (e netTimeoutError) Error() string { return "i/o timeout" }
|
||||||
|
func (e netTimeoutError) Timeout() bool { return true }
|
||||||
|
func (e netTimeoutError) Temporary() bool { return true }
|
||||||
|
|
||||||
type ConnRewindBaseConn struct {
|
type ConnRewindBaseConn struct {
|
||||||
testlib.EssentialsConnMock
|
testlib.EssentialsConnMock
|
||||||
|
|
||||||
@@ -291,6 +297,141 @@ func (suite *ConnProxyProtocolTestSuite) TearDownTest() {
|
|||||||
suite.targetConnMock.AssertExpectations(suite.T())
|
suite.targetConnMock.AssertExpectations(suite.T())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type IdleTrackerTestSuite struct {
|
||||||
|
suite.Suite
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *IdleTrackerTestSuite) TestNewNotIdle() {
|
||||||
|
tracker := newIdleTracker(time.Second)
|
||||||
|
suite.False(tracker.isIdle())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *IdleTrackerTestSuite) TestIdleAfterTimeout() {
|
||||||
|
tracker := newIdleTracker(10 * time.Millisecond)
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
|
||||||
|
suite.True(tracker.isIdle())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *IdleTrackerTestSuite) TestTouchResetsIdle() {
|
||||||
|
tracker := newIdleTracker(50 * time.Millisecond)
|
||||||
|
time.Sleep(30 * time.Millisecond)
|
||||||
|
|
||||||
|
tracker.touch()
|
||||||
|
|
||||||
|
suite.False(tracker.isIdle())
|
||||||
|
}
|
||||||
|
|
||||||
|
type ConnIdleTimeoutTestSuite struct {
|
||||||
|
suite.Suite
|
||||||
|
|
||||||
|
connMock *testlib.EssentialsConnMock
|
||||||
|
tracker *idleTracker
|
||||||
|
conn connIdleTimeout
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) SetupTest() {
|
||||||
|
suite.connMock = &testlib.EssentialsConnMock{}
|
||||||
|
suite.tracker = newIdleTracker(time.Second)
|
||||||
|
suite.conn = connIdleTimeout{
|
||||||
|
Conn: suite.connMock,
|
||||||
|
tracker: suite.tracker,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TearDownTest() {
|
||||||
|
suite.connMock.AssertExpectations(suite.T())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TestReadOk() {
|
||||||
|
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
|
||||||
|
suite.connMock.On("Read", mock.Anything).Once().Return(5, nil)
|
||||||
|
|
||||||
|
n, err := suite.conn.Read(make([]byte, 10))
|
||||||
|
suite.NoError(err)
|
||||||
|
suite.Equal(5, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TestReadNonTimeoutErr() {
|
||||||
|
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
|
||||||
|
suite.connMock.On("Read", mock.Anything).Once().Return(0, io.EOF)
|
||||||
|
|
||||||
|
n, err := suite.conn.Read(make([]byte, 10))
|
||||||
|
suite.True(errors.Is(err, io.EOF))
|
||||||
|
suite.Equal(0, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TestReadTimeoutRetriesWhenNotIdle() {
|
||||||
|
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
|
||||||
|
suite.connMock.On("Read", mock.Anything).Once().Return(0, netTimeoutError{})
|
||||||
|
suite.connMock.On("Read", mock.Anything).Once().Return(5, nil)
|
||||||
|
|
||||||
|
n, err := suite.conn.Read(make([]byte, 10))
|
||||||
|
suite.NoError(err)
|
||||||
|
suite.Equal(5, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TestReadTimeoutClosesWhenIdle() {
|
||||||
|
suite.tracker = newIdleTracker(time.Millisecond)
|
||||||
|
suite.conn = connIdleTimeout{
|
||||||
|
Conn: suite.connMock,
|
||||||
|
tracker: suite.tracker,
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(5 * time.Millisecond)
|
||||||
|
|
||||||
|
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
|
||||||
|
suite.connMock.On("Read", mock.Anything).Once().Return(0, netTimeoutError{})
|
||||||
|
|
||||||
|
n, err := suite.conn.Read(make([]byte, 10))
|
||||||
|
suite.Equal(0, n)
|
||||||
|
|
||||||
|
netErr, ok := err.(net.Error) //nolint: errorlint
|
||||||
|
suite.True(ok)
|
||||||
|
suite.True(netErr.Timeout())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TestSharedTrackerPreventsFalseTimeout() {
|
||||||
|
connMock2 := &testlib.EssentialsConnMock{}
|
||||||
|
conn2 := connIdleTimeout{
|
||||||
|
Conn: connMock2,
|
||||||
|
tracker: suite.tracker,
|
||||||
|
}
|
||||||
|
|
||||||
|
connMock2.On("SetWriteDeadline", mock.Anything).Return(nil)
|
||||||
|
connMock2.On("Write", mock.Anything).Once().Return(5, nil)
|
||||||
|
|
||||||
|
_, _ = conn2.Write(make([]byte, 5))
|
||||||
|
|
||||||
|
suite.connMock.On("SetReadDeadline", mock.Anything).Return(nil)
|
||||||
|
suite.connMock.On("Read", mock.Anything).Once().Return(0, netTimeoutError{})
|
||||||
|
suite.connMock.On("Read", mock.Anything).Once().Return(3, nil)
|
||||||
|
|
||||||
|
n, err := suite.conn.Read(make([]byte, 10))
|
||||||
|
suite.NoError(err)
|
||||||
|
suite.Equal(3, n)
|
||||||
|
|
||||||
|
connMock2.AssertExpectations(suite.T())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TestWriteOk() {
|
||||||
|
suite.connMock.On("SetWriteDeadline", mock.Anything).Return(nil)
|
||||||
|
suite.connMock.On("Write", mock.Anything).Once().Return(5, nil)
|
||||||
|
|
||||||
|
n, err := suite.conn.Write(make([]byte, 5))
|
||||||
|
suite.NoError(err)
|
||||||
|
suite.Equal(5, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (suite *ConnIdleTimeoutTestSuite) TestWriteErr() {
|
||||||
|
suite.connMock.On("SetWriteDeadline", mock.Anything).Return(nil)
|
||||||
|
suite.connMock.On("Write", mock.Anything).Once().Return(0, io.EOF)
|
||||||
|
|
||||||
|
n, err := suite.conn.Write(make([]byte, 5))
|
||||||
|
suite.True(errors.Is(err, io.EOF))
|
||||||
|
suite.Equal(0, n)
|
||||||
|
}
|
||||||
|
|
||||||
func TestConnTraffic(t *testing.T) {
|
func TestConnTraffic(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
suite.Run(t, &ConnTrafficTestSuite{})
|
suite.Run(t, &ConnTrafficTestSuite{})
|
||||||
@@ -305,3 +446,13 @@ func TestConnProxyProtocol(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
suite.Run(t, &ConnProxyProtocolTestSuite{})
|
suite.Run(t, &ConnProxyProtocolTestSuite{})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestIdleTracker(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
suite.Run(t, &IdleTrackerTestSuite{})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConnIdleTimeout(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
suite.Run(t, &ConnIdleTimeoutTestSuite{})
|
||||||
|
}
|
||||||
|
|||||||
+8
-4
@@ -102,11 +102,13 @@ func (p *Proxy) ServeConn(conn essentials.Conn) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
tracker := newIdleTracker(p.idleTimeout)
|
||||||
|
|
||||||
relay.Relay(
|
relay.Relay(
|
||||||
ctx,
|
ctx,
|
||||||
ctx.logger.Named("relay"),
|
ctx.logger.Named("relay"),
|
||||||
connIdleTimeout{Conn: ctx.telegramConn, timeout: p.idleTimeout},
|
connIdleTimeout{Conn: ctx.telegramConn, tracker: tracker},
|
||||||
connIdleTimeout{Conn: ctx.clientConn, timeout: p.idleTimeout},
|
connIdleTimeout{Conn: ctx.clientConn, tracker: tracker},
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -305,11 +307,13 @@ func (p *Proxy) doDomainFronting(ctx *streamContext, conn *connRewind) {
|
|||||||
stream: p.eventStream,
|
stream: p.eventStream,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
tracker := newIdleTracker(p.idleTimeout)
|
||||||
|
|
||||||
relay.Relay(
|
relay.Relay(
|
||||||
ctx,
|
ctx,
|
||||||
ctx.logger.Named("domain-fronting"),
|
ctx.logger.Named("domain-fronting"),
|
||||||
connIdleTimeout{Conn: frontConn, timeout: p.idleTimeout},
|
connIdleTimeout{Conn: frontConn, tracker: tracker},
|
||||||
connIdleTimeout{Conn: conn, timeout: p.idleTimeout},
|
connIdleTimeout{Conn: conn, tracker: tracker},
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user