Merge remote-tracking branch 'origin/stable' into v2

This commit is contained in:
9seconds
2026-03-30 13:09:05 +02:00
4 changed files with 212 additions and 9 deletions
BIN
View File
Binary file not shown.
+53 -5
View File
@@ -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
} }
+151
View File
@@ -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
View File
@@ -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},
) )
} }