mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 13:34:02 +03:00
Merge remote-tracking branch 'origin/stable' into v2
This commit is contained in:
BIN
Binary file not shown.
+53
-5
@@ -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
|
||||
}
|
||||
|
||||
@@ -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{})
|
||||
}
|
||||
|
||||
+8
-4
@@ -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},
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user