diff --git a/mtglib/conns.go b/mtglib/conns.go index 12f19d1..55d57ee 100644 --- a/mtglib/conns.go +++ b/mtglib/conns.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "net" + "sync/atomic" "time" "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 { essentials.Conn - timeout time.Duration + tracker *idleTracker } 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) { - 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 } diff --git a/mtglib/conns_internal_test.go b/mtglib/conns_internal_test.go index 123258f..21f56ee 100644 --- a/mtglib/conns_internal_test.go +++ b/mtglib/conns_internal_test.go @@ -16,6 +16,12 @@ import ( "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 { testlib.EssentialsConnMock @@ -291,6 +297,141 @@ func (suite *ConnProxyProtocolTestSuite) TearDownTest() { 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) { t.Parallel() suite.Run(t, &ConnTrafficTestSuite{}) @@ -305,3 +446,13 @@ func TestConnProxyProtocol(t *testing.T) { t.Parallel() 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{}) +} diff --git a/mtglib/proxy.go b/mtglib/proxy.go index 8a904ee..7acf8b7 100644 --- a/mtglib/proxy.go +++ b/mtglib/proxy.go @@ -102,11 +102,13 @@ func (p *Proxy) ServeConn(conn essentials.Conn) { return } + tracker := newIdleTracker(p.idleTimeout) + relay.Relay( ctx, ctx.logger.Named("relay"), - connIdleTimeout{Conn: ctx.telegramConn, timeout: p.idleTimeout}, - connIdleTimeout{Conn: ctx.clientConn, timeout: p.idleTimeout}, + connIdleTimeout{Conn: ctx.telegramConn, tracker: tracker}, + connIdleTimeout{Conn: ctx.clientConn, tracker: tracker}, ) } @@ -305,11 +307,13 @@ func (p *Proxy) doDomainFronting(ctx *streamContext, conn *connRewind) { stream: p.eventStream, } + tracker := newIdleTracker(p.idleTimeout) + relay.Relay( ctx, ctx.logger.Named("domain-fronting"), - connIdleTimeout{Conn: frontConn, timeout: p.idleTimeout}, - connIdleTimeout{Conn: conn, timeout: p.idleTimeout}, + connIdleTimeout{Conn: frontConn, tracker: tracker}, + connIdleTimeout{Conn: conn, tracker: tracker}, ) }