Use CloseRead and CloseWrites

This commit is contained in:
9seconds
2021-12-01 10:37:31 +03:00
parent 7b1f86b75d
commit ffad717829
32 changed files with 320 additions and 99 deletions
+17
View File
@@ -0,0 +1,17 @@
package essentials
import "net"
type CloseableReader interface {
CloseRead() error
}
type CloseableWriter interface {
CloseWrite() error
}
type Conn interface {
net.Conn
CloseableReader
CloseableWriter
}
+4
View File
@@ -38,9 +38,13 @@ require (
github.com/gotd/ige v0.1.5 // indirect github.com/gotd/ige v0.1.5 // indirect
github.com/gotd/xor v0.1.1 // indirect github.com/gotd/xor v0.1.1 // indirect
github.com/matttproud/golang_protobuf_extensions v1.0.1 // indirect github.com/matttproud/golang_protobuf_extensions v1.0.1 // indirect
github.com/patrickmn/go-cache v2.1.0+incompatible // indirect
github.com/pkg/errors v0.9.1 // indirect github.com/pkg/errors v0.9.1 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/prometheus/client_model v0.2.0 // indirect github.com/prometheus/client_model v0.2.0 // indirect
github.com/txthinking/runnergroup v0.0.0-20210608031112-152c7c4432bf // indirect
github.com/txthinking/socks5 v0.0.0-20211121111206-e03c1217a50b // indirect
github.com/txthinking/x v0.0.0-20210326105829-476fab902fbe // indirect
go.uber.org/atomic v1.7.0 // indirect go.uber.org/atomic v1.7.0 // indirect
go.uber.org/multierr v1.6.0 // indirect go.uber.org/multierr v1.6.0 // indirect
go.uber.org/zap v1.16.0 // indirect go.uber.org/zap v1.16.0 // indirect
+8
View File
@@ -195,6 +195,8 @@ github.com/mwitkow/go-conntrack v0.0.0-20161129095857-cc309e4a2223/go.mod h1:qRW
github.com/mwitkow/go-conntrack v0.0.0-20190716064945-2f068394615f/go.mod h1:qRWi+5nqEBWmkhHvq77mSJWrCKwh8bxhgT7d/eI7P4U= github.com/mwitkow/go-conntrack v0.0.0-20190716064945-2f068394615f/go.mod h1:qRWi+5nqEBWmkhHvq77mSJWrCKwh8bxhgT7d/eI7P4U=
github.com/panjf2000/ants/v2 v2.4.6 h1:drmj9mcygn2gawZ155dRbo+NfXEfAssjZNU1qoIb4gQ= github.com/panjf2000/ants/v2 v2.4.6 h1:drmj9mcygn2gawZ155dRbo+NfXEfAssjZNU1qoIb4gQ=
github.com/panjf2000/ants/v2 v2.4.6/go.mod h1:f6F0NZVFsGCp5A7QW/Zj/m92atWwOkY0OIhFxRNFr4A= github.com/panjf2000/ants/v2 v2.4.6/go.mod h1:f6F0NZVFsGCp5A7QW/Zj/m92atWwOkY0OIhFxRNFr4A=
github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc=
github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ=
github.com/pelletier/go-toml v1.9.4 h1:tjENF6MfZAg8e4ZmZTeWaWiT2vXtsoO6+iuOjFhECwM= github.com/pelletier/go-toml v1.9.4 h1:tjENF6MfZAg8e4ZmZTeWaWiT2vXtsoO6+iuOjFhECwM=
github.com/pelletier/go-toml v1.9.4/go.mod h1:u1nR/EPcESfeI/szUZKdtJ0xRNbUoANCkoOuaOx1Y+c= github.com/pelletier/go-toml v1.9.4/go.mod h1:u1nR/EPcESfeI/szUZKdtJ0xRNbUoANCkoOuaOx1Y+c=
github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA= github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA=
@@ -250,6 +252,12 @@ github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81P
github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.0 h1:nwc3DEeHmmLAfoZucVR881uASk0Mfjw8xYJ99tb5CcY= github.com/stretchr/testify v1.7.0 h1:nwc3DEeHmmLAfoZucVR881uASk0Mfjw8xYJ99tb5CcY=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/txthinking/runnergroup v0.0.0-20210608031112-152c7c4432bf h1:7PflaKRtU4np/epFxRXlFhlzLXZzKFrH5/I4so5Ove0=
github.com/txthinking/runnergroup v0.0.0-20210608031112-152c7c4432bf/go.mod h1:CLUSJbazqETbaR+i0YAhXBICV9TrKH93pziccMhmhpM=
github.com/txthinking/socks5 v0.0.0-20211121111206-e03c1217a50b h1:6J/38A0Xmdnjacfie0Udams7OP/GdoExyTipKwuQWjY=
github.com/txthinking/socks5 v0.0.0-20211121111206-e03c1217a50b/go.mod h1:7NloQcrxaZYKURWph5HLxVDlIwMHJXCPkeWPtpftsIg=
github.com/txthinking/x v0.0.0-20210326105829-476fab902fbe h1:gMWxZxBFRAXqoGkwkYlPX2zvyyKNWJpxOxCrjqJkm5A=
github.com/txthinking/x v0.0.0-20210326105829-476fab902fbe/go.mod h1:WgqbSEmUYSjEV3B1qmee/PpP2NYEz4bL9/+mF1ma+s4=
github.com/tylertreat/BoomFilters v0.0.0-20210315201527-1a82519a3e43 h1:QEePdg0ty2r0t1+qwfZmQ4OOl/MB2UXIeJSpIZv56lg= github.com/tylertreat/BoomFilters v0.0.0-20210315201527-1a82519a3e43 h1:QEePdg0ty2r0t1+qwfZmQ4OOl/MB2UXIeJSpIZv56lg=
github.com/tylertreat/BoomFilters v0.0.0-20210315201527-1a82519a3e43/go.mod h1:OYRfF6eb5wY9VRFkXJH8FFBi3plw2v+giaIu7P054pM= github.com/tylertreat/BoomFilters v0.0.0-20210315201527-1a82519a3e43/go.mod h1:OYRfF6eb5wY9VRFkXJH8FFBi3plw2v+giaIu7P054pM=
github.com/yuin/goldmark v1.1.25/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= github.com/yuin/goldmark v1.1.25/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
+2 -1
View File
@@ -13,6 +13,7 @@ import (
"strings" "strings"
"sync" "sync"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/internal/config" "github.com/9seconds/mtg/v2/internal/config"
"github.com/9seconds/mtg/v2/internal/utils" "github.com/9seconds/mtg/v2/internal/utils"
"github.com/9seconds/mtg/v2/mtglib" "github.com/9seconds/mtg/v2/mtglib"
@@ -106,7 +107,7 @@ func (a *Access) Run(cli *CLI, version string) error {
} }
func (a *Access) getIP(ntw mtglib.Network, protocol string) net.IP { func (a *Access) getIP(ntw mtglib.Network, protocol string) net.IP {
client := ntw.MakeHTTPClient(func(ctx context.Context, network, address string) (net.Conn, error) { client := ntw.MakeHTTPClient(func(ctx context.Context, network, address string) (essentials.Conn, error) {
return ntw.DialContext(ctx, protocol, address) // nolint: wrapcheck return ntw.DialContext(ctx, protocol, address) // nolint: wrapcheck
}) })
+6 -6
View File
@@ -2,9 +2,9 @@ package testlib
import ( import (
"context" "context"
"net"
"net/http" "net/http"
"github.com/9seconds/mtg/v2/essentials"
"github.com/stretchr/testify/mock" "github.com/stretchr/testify/mock"
) )
@@ -12,19 +12,19 @@ type MtglibNetworkMock struct {
mock.Mock mock.Mock
} }
func (m *MtglibNetworkMock) Dial(network, address string) (net.Conn, error) { func (m *MtglibNetworkMock) Dial(network, address string) (essentials.Conn, error) {
args := m.Called(network, address) args := m.Called(network, address)
return args.Get(0).(net.Conn), args.Error(1) // nolint: wrapcheck return args.Get(0).(essentials.Conn), args.Error(1) // nolint: wrapcheck
} }
func (m *MtglibNetworkMock) DialContext(ctx context.Context, network, address string) (net.Conn, error) { func (m *MtglibNetworkMock) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) {
args := m.Called(ctx, network, address) args := m.Called(ctx, network, address)
return args.Get(0).(net.Conn), args.Error(1) // nolint: wrapcheck return args.Get(0).(essentials.Conn), args.Error(1) // nolint: wrapcheck
} }
func (m *MtglibNetworkMock) MakeHTTPClient(dialFunc func(ctx context.Context, func (m *MtglibNetworkMock) MakeHTTPClient(dialFunc func(ctx context.Context,
network, address string) (net.Conn, error)) *http.Client { network, address string) (essentials.Conn, error)) *http.Client {
return m.Called(dialFunc).Get(0).(*http.Client) return m.Called(dialFunc).Get(0).(*http.Client)
} }
+17 -9
View File
@@ -7,42 +7,50 @@ import (
"github.com/stretchr/testify/mock" "github.com/stretchr/testify/mock"
) )
type NetConnMock struct { type EssentialsConnMock struct {
mock.Mock mock.Mock
} }
func (n *NetConnMock) Read(b []byte) (int, error) { func (n *EssentialsConnMock) Read(b []byte) (int, error) {
args := n.Called(b) args := n.Called(b)
return args.Int(0), args.Error(1) return args.Int(0), args.Error(1)
} }
func (n *NetConnMock) Write(b []byte) (int, error) { func (n *EssentialsConnMock) Write(b []byte) (int, error) {
args := n.Called(b) args := n.Called(b)
return args.Int(0), args.Error(1) return args.Int(0), args.Error(1)
} }
func (n *NetConnMock) Close() error { func (n *EssentialsConnMock) Close() error {
return n.Called().Error(0) // nolint: wrapcheck return n.Called().Error(0) // nolint: wrapcheck
} }
func (n *NetConnMock) LocalAddr() net.Addr { func (n *EssentialsConnMock) CloseRead() error {
return n.Called().Error(0) // nolint: wrapcheck
}
func (n *EssentialsConnMock) CloseWrite() error {
return n.Called().Error(0) // nolint: wrapcheck
}
func (n *EssentialsConnMock) LocalAddr() net.Addr {
return n.Called().Get(0).(net.Addr) return n.Called().Get(0).(net.Addr)
} }
func (n *NetConnMock) RemoteAddr() net.Addr { func (n *EssentialsConnMock) RemoteAddr() net.Addr {
return n.Called().Get(0).(net.Addr) return n.Called().Get(0).(net.Addr)
} }
func (n *NetConnMock) SetDeadline(t time.Time) error { func (n *EssentialsConnMock) SetDeadline(t time.Time) error {
return n.Called(t).Error(0) // nolint: wrapcheck return n.Called(t).Error(0) // nolint: wrapcheck
} }
func (n *NetConnMock) SetReadDeadline(t time.Time) error { func (n *EssentialsConnMock) SetReadDeadline(t time.Time) error {
return n.Called(t).Error(0) // nolint: wrapcheck return n.Called(t).Error(0) // nolint: wrapcheck
} }
func (n *NetConnMock) SetWriteDeadline(t time.Time) error { func (n *EssentialsConnMock) SetWriteDeadline(t time.Time) error {
return n.Called(t).Error(0) // nolint: wrapcheck return n.Called(t).Error(0) // nolint: wrapcheck
} }
+5 -4
View File
@@ -4,12 +4,13 @@ import (
"bytes" "bytes"
"context" "context"
"io" "io"
"net"
"sync" "sync"
"github.com/9seconds/mtg/v2/essentials"
) )
type connTraffic struct { type connTraffic struct {
net.Conn essentials.Conn
streamID string streamID string
stream EventStream stream EventStream
@@ -37,7 +38,7 @@ func (c connTraffic) Write(b []byte) (int, error) {
} }
type connRewind struct { type connRewind struct {
net.Conn essentials.Conn
active io.Reader active io.Reader
buf bytes.Buffer buf bytes.Buffer
@@ -58,7 +59,7 @@ func (c *connRewind) Rewind() {
c.active = io.MultiReader(&c.buf, c.Conn) c.active = io.MultiReader(&c.buf, c.Conn)
} }
func newConnRewind(conn net.Conn) *connRewind { func newConnRewind(conn essentials.Conn) *connRewind {
rv := &connRewind{ rv := &connRewind{
Conn: conn, Conn: conn,
} }
+3 -3
View File
@@ -14,7 +14,7 @@ import (
) )
type ConnRewindBaseConn struct { type ConnRewindBaseConn struct {
testlib.NetConnMock testlib.EssentialsConnMock
readBuffer bytes.Buffer readBuffer bytes.Buffer
} }
@@ -29,13 +29,13 @@ type ConnTrafficTestSuite struct {
suite.Suite suite.Suite
eventStreamMock *EventStreamMock eventStreamMock *EventStreamMock
connMock *testlib.NetConnMock connMock *testlib.EssentialsConnMock
conn io.ReadWriter conn io.ReadWriter
} }
func (suite *ConnTrafficTestSuite) SetupTest() { func (suite *ConnTrafficTestSuite) SetupTest() {
suite.eventStreamMock = &EventStreamMock{} suite.eventStreamMock = &EventStreamMock{}
suite.connMock = &testlib.NetConnMock{} suite.connMock = &testlib.EssentialsConnMock{}
suite.conn = connTraffic{ suite.conn = connTraffic{
Conn: suite.connMock, Conn: suite.connMock,
streamID: "CONNID", streamID: "CONNID",
+5 -3
View File
@@ -23,6 +23,8 @@ import (
"net" "net"
"net/http" "net/http"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
) )
var ( var (
@@ -116,16 +118,16 @@ const (
// 3. Doing HTTP requests (for example, for FireHOL ipblocklist). // 3. Doing HTTP requests (for example, for FireHOL ipblocklist).
type Network interface { type Network interface {
// Dial establishes context-free TCP connections. // Dial establishes context-free TCP connections.
Dial(network, address string) (net.Conn, error) Dial(network, address string) (essentials.Conn, error)
// DialContext dials using a context. This is a preferrable // DialContext dials using a context. This is a preferrable
// way of establishing TCP connections. // way of establishing TCP connections.
DialContext(ctx context.Context, network, address string) (net.Conn, error) DialContext(ctx context.Context, network, address string) (essentials.Conn, error)
// MakeHTTPClient build an HTTP client with given dial function. If // MakeHTTPClient build an HTTP client with given dial function. If
// nothing is provided, then DialContext of this interface is going // nothing is provided, then DialContext of this interface is going
// to be used. // to be used.
MakeHTTPClient(func(ctx context.Context, network, address string) (net.Conn, error)) *http.Client MakeHTTPClient(func(ctx context.Context, network, address string) (essentials.Conn, error)) *http.Client
} }
// AntiReplayCache is an interface that is used to detect replay attacks // AntiReplayCache is an interface that is used to detect replay attacks
+2 -2
View File
@@ -4,13 +4,13 @@ import (
"bytes" "bytes"
"fmt" "fmt"
"math/rand" "math/rand"
"net"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/mtglib/internal/faketls/record" "github.com/9seconds/mtg/v2/mtglib/internal/faketls/record"
) )
type Conn struct { type Conn struct {
net.Conn essentials.Conn
readBuffer bytes.Buffer readBuffer bytes.Buffer
} }
+1 -1
View File
@@ -15,7 +15,7 @@ import (
) )
type ConnMock struct { type ConnMock struct {
testlib.NetConnMock testlib.EssentialsConnMock
readBuffer bytes.Buffer readBuffer bytes.Buffer
writeBuffer bytes.Buffer writeBuffer bytes.Buffer
@@ -42,7 +42,7 @@ func (suite *ClientHandshakeTestSuite) TestOk() {
writeData := make([]byte, len(snapshot.Encrypted.Text.data)) writeData := make([]byte, len(snapshot.Encrypted.Text.data))
readData := make([]byte, len(snapshot.Decrypted.Text.data)) readData := make([]byte, len(snapshot.Decrypted.Text.data))
connMock := &testlib.NetConnMock{} connMock := &testlib.EssentialsConnMock{}
connMock.On("Read", mock.Anything). connMock.On("Read", mock.Anything).
Once(). Once().
Return(len(snapshot.Decrypted.Text.data), nil). Return(len(snapshot.Decrypted.Text.data), nil).
+3 -2
View File
@@ -2,11 +2,12 @@ package obfuscated2
import ( import (
"crypto/cipher" "crypto/cipher"
"net"
"github.com/9seconds/mtg/v2/essentials"
) )
type Conn struct { type Conn struct {
net.Conn essentials.Conn
Encryptor cipher.Stream Encryptor cipher.Stream
Decryptor cipher.Stream Decryptor cipher.Stream
@@ -16,7 +16,7 @@ import (
type ServerHandshakeTestSuite struct { type ServerHandshakeTestSuite struct {
suite.Suite suite.Suite
connMock *testlib.NetConnMock connMock *testlib.EssentialsConnMock
proxyConn obfuscated2.Conn proxyConn obfuscated2.Conn
encryptor cipher.Stream encryptor cipher.Stream
decryptor cipher.Stream decryptor cipher.Stream
@@ -24,7 +24,7 @@ type ServerHandshakeTestSuite struct {
func (suite *ServerHandshakeTestSuite) SetupTest() { func (suite *ServerHandshakeTestSuite) SetupTest() {
buf := &bytes.Buffer{} buf := &bytes.Buffer{}
suite.connMock = &testlib.NetConnMock{} suite.connMock = &testlib.EssentialsConnMock{}
encryptor, decryptor, err := obfuscated2.ServerHandshake(buf) encryptor, decryptor, err := obfuscated2.ServerHandshake(buf)
suite.NoError(err) suite.NoError(err)
+23 -9
View File
@@ -2,11 +2,14 @@ package relay
import ( import (
"context" "context"
"net" "errors"
"io"
"sync" "sync"
"github.com/9seconds/mtg/v2/essentials"
) )
func Relay(ctx context.Context, log Logger, telegramConn, clientConn net.Conn) { func Relay(ctx context.Context, log Logger, telegramConn, clientConn essentials.Conn) {
defer telegramConn.Close() defer telegramConn.Close()
defer clientConn.Close() defer clientConn.Close()
@@ -29,14 +32,25 @@ func Relay(ctx context.Context, log Logger, telegramConn, clientConn net.Conn) {
wg.Wait() wg.Wait()
} }
func pump(log Logger, src, dst net.Conn, wg *sync.WaitGroup, direction string) { func pump(log Logger, src, dst essentials.Conn, wg *sync.WaitGroup, direction string) {
defer wg.Done()
syncer := acquireSyncPair(src, dst) syncer := acquireSyncPair(src, dst)
defer releaseSyncPair(syncer)
defer syncer.Flush()
if n, err := syncer.Sync(); err != nil { defer func() {
log.Printf("cannot pump %s (written %d bytes): %v", direction, n, err) syncer.Flush()
releaseSyncPair(syncer)
src.CloseRead()
dst.CloseWrite()
wg.Done()
}()
n, err := syncer.Sync()
switch {
case err == nil:
log.Printf("%s has been finished", direction)
case errors.Is(err, io.EOF):
log.Printf("%s has been finished because of EOF. Written %d bytes", direction, n)
default:
log.Printf("%s has been finished (written %d bytes): %v", direction, n, err)
} }
} }
+10 -4
View File
@@ -17,8 +17,8 @@ type RelayTestSuite struct {
loggerMock relay.Logger loggerMock relay.Logger
ctx context.Context ctx context.Context
ctxCancel context.CancelFunc ctxCancel context.CancelFunc
telegramConnMock *testlib.NetConnMock telegramConnMock *testlib.EssentialsConnMock
clientConnMock *testlib.NetConnMock clientConnMock *testlib.EssentialsConnMock
} }
func (suite *RelayTestSuite) SetupTest() { func (suite *RelayTestSuite) SetupTest() {
@@ -26,8 +26,8 @@ func (suite *RelayTestSuite) SetupTest() {
suite.ctx = ctx suite.ctx = ctx
suite.ctxCancel = cancel suite.ctxCancel = cancel
suite.loggerMock = &loggerMock{} suite.loggerMock = &loggerMock{}
suite.telegramConnMock = &testlib.NetConnMock{} suite.telegramConnMock = &testlib.EssentialsConnMock{}
suite.clientConnMock = &testlib.NetConnMock{} suite.clientConnMock = &testlib.EssentialsConnMock{}
} }
func (suite *RelayTestSuite) TearDownTest() { func (suite *RelayTestSuite) TearDownTest() {
@@ -38,12 +38,18 @@ func (suite *RelayTestSuite) TearDownTest() {
func (suite *RelayTestSuite) TestExit() { func (suite *RelayTestSuite) TestExit() {
suite.telegramConnMock.On("Close").Return(nil) suite.telegramConnMock.On("Close").Return(nil)
suite.telegramConnMock.On("CloseRead").Return(nil).Once()
suite.telegramConnMock.On("CloseWrite").Return(nil).Once()
suite.telegramConnMock.On("Read", mock.Anything).Return(10, io.EOF).Once() suite.telegramConnMock.On("Read", mock.Anything).Return(10, io.EOF).Once()
suite.telegramConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe() suite.telegramConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
suite.telegramConnMock.On("SetReadDeadline", mock.Anything).Return(nil).Maybe()
suite.clientConnMock.On("Read", mock.Anything).Return(0, io.EOF).Once() suite.clientConnMock.On("Read", mock.Anything).Return(0, io.EOF).Once()
suite.clientConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe() suite.clientConnMock.On("Write", mock.Anything).Return(10, io.EOF).Maybe()
suite.clientConnMock.On("Close").Return(nil) suite.clientConnMock.On("Close").Return(nil)
suite.clientConnMock.On("CloseRead").Return(nil).Once()
suite.clientConnMock.On("CloseWrite").Return(nil).Once()
suite.clientConnMock.On("SetReadDeadline", mock.Anything).Return(nil).Maybe()
relay.Relay(suite.ctx, suite.loggerMock, suite.telegramConnMock, suite.clientConnMock) relay.Relay(suite.ctx, suite.loggerMock, suite.telegramConnMock, suite.clientConnMock)
} }
+1 -1
View File
@@ -48,7 +48,7 @@ func (s *syncPair) Flush() error {
s.mutex.Lock() s.mutex.Lock()
defer s.mutex.Unlock() defer s.mutex.Unlock()
return s.writer.Flush() return s.writer.Flush() // nolint: wrapcheck
} }
func (s *syncPair) readBlocking(p []byte, blocking bool) (int, error) { func (s *syncPair) readBlocking(p []byte, blocking bool) (int, error) {
+3 -2
View File
@@ -2,7 +2,8 @@ package telegram
import ( import (
"context" "context"
"net"
"github.com/9seconds/mtg/v2/essentials"
) )
type preferIP uint8 type preferIP uint8
@@ -82,5 +83,5 @@ var (
) )
type Dialer interface { type Dialer interface {
DialContext(ctx context.Context, network, address string) (net.Conn, error) DialContext(ctx context.Context, network, address string) (essentials.Conn, error)
} }
+4 -3
View File
@@ -3,8 +3,9 @@ package telegram
import ( import (
"context" "context"
"fmt" "fmt"
"net"
"strings" "strings"
"github.com/9seconds/mtg/v2/essentials"
) )
type Telegram struct { type Telegram struct {
@@ -13,7 +14,7 @@ type Telegram struct {
pool addressPool pool addressPool
} }
func (t Telegram) Dial(ctx context.Context, dc int) (net.Conn, error) { func (t Telegram) Dial(ctx context.Context, dc int) (essentials.Conn, error) {
var addresses []tgAddr var addresses []tgAddr
switch t.preferIP { switch t.preferIP {
@@ -28,7 +29,7 @@ func (t Telegram) Dial(ctx context.Context, dc int) (net.Conn, error) {
} }
var ( var (
conn net.Conn conn essentials.Conn
err error err error
) )
+3 -2
View File
@@ -9,6 +9,7 @@ import (
"sync" "sync"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/mtglib/internal/faketls" "github.com/9seconds/mtg/v2/mtglib/internal/faketls"
"github.com/9seconds/mtg/v2/mtglib/internal/faketls/record" "github.com/9seconds/mtg/v2/mtglib/internal/faketls/record"
"github.com/9seconds/mtg/v2/mtglib/internal/obfuscated2" "github.com/9seconds/mtg/v2/mtglib/internal/obfuscated2"
@@ -44,7 +45,7 @@ func (p *Proxy) DomainFrontingAddress() string {
// ServeConn serves a connection. We do not check IP blocklist and // ServeConn serves a connection. We do not check IP blocklist and
// concurrency limit here. // concurrency limit here.
func (p *Proxy) ServeConn(conn net.Conn) { func (p *Proxy) ServeConn(conn essentials.Conn) {
p.streamWaitGroup.Add(1) p.streamWaitGroup.Add(1)
defer p.streamWaitGroup.Done() defer p.streamWaitGroup.Done()
@@ -299,7 +300,7 @@ func NewProxy(opts ProxyOpts) (*Proxy, error) {
pool, err := ants.NewPoolWithFunc(opts.getConcurrency(), pool, err := ants.NewPoolWithFunc(opts.getConcurrency(),
func(arg interface{}) { func(arg interface{}) {
proxy.ServeConn(arg.(net.Conn)) proxy.ServeConn(arg.(essentials.Conn))
}, },
ants.WithLogger(opts.getLogger("ants")), ants.WithLogger(opts.getLogger("ants")),
ants.WithNonblocking(true)) ants.WithNonblocking(true))
+5 -3
View File
@@ -6,13 +6,15 @@ import (
"encoding/base64" "encoding/base64"
"net" "net"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
) )
type streamContext struct { type streamContext struct {
ctx context.Context ctx context.Context
ctxCancel context.CancelFunc ctxCancel context.CancelFunc
clientConn net.Conn clientConn essentials.Conn
telegramConn net.Conn telegramConn essentials.Conn
streamID string streamID string
dc int dc int
logger Logger logger Logger
@@ -50,7 +52,7 @@ func (s *streamContext) ClientIP() net.IP {
return s.clientConn.RemoteAddr().(*net.TCPAddr).IP return s.clientConn.RemoteAddr().(*net.TCPAddr).IP
} }
func newStreamContext(ctx context.Context, logger Logger, clientConn net.Conn) *streamContext { func newStreamContext(ctx context.Context, logger Logger, clientConn essentials.Conn) *streamContext {
connIDBytes := make([]byte, ConnectionIDBytesLength) connIDBytes := make([]byte, ConnectionIDBytesLength)
if _, err := rand.Read(connIDBytes); err != nil { if _, err := rand.Read(connIDBytes); err != nil {
+3 -3
View File
@@ -12,7 +12,7 @@ import (
type StreamContextTestSuite struct { type StreamContextTestSuite struct {
suite.Suite suite.Suite
connMock *testlib.NetConnMock connMock *testlib.EssentialsConnMock
logger NoopLogger logger NoopLogger
ctx *streamContext ctx *streamContext
ctxCancel context.CancelFunc ctxCancel context.CancelFunc
@@ -27,7 +27,7 @@ func (suite *StreamContextTestSuite) SetupTest() {
ctx = context.WithValue(ctx, "key", "value") // nolint: golint, revive, staticcheck ctx = context.WithValue(ctx, "key", "value") // nolint: golint, revive, staticcheck
suite.ctxCancel = cancel suite.ctxCancel = cancel
suite.connMock = &testlib.NetConnMock{} suite.connMock = &testlib.EssentialsConnMock{}
addr := &net.TCPAddr{ addr := &net.TCPAddr{
IP: net.ParseIP("10.0.0.10"), IP: net.ParseIP("10.0.0.10"),
@@ -73,7 +73,7 @@ func (suite *StreamContextTestSuite) TestClientIP() {
func (suite *StreamContextTestSuite) TestClose() { func (suite *StreamContextTestSuite) TestClose() {
suite.connMock.On("Close").Once().Return(nil) suite.connMock.On("Close").Once().Return(nil)
tgConnMock := &testlib.NetConnMock{} tgConnMock := &testlib.EssentialsConnMock{}
tgConnMock.On("Close").Once().Return(nil) tgConnMock.On("Close").Once().Return(nil)
suite.ctx.telegramConn = tgConnMock suite.ctx.telegramConn = tgConnMock
+7 -5
View File
@@ -2,9 +2,10 @@ package network
import ( import (
"context" "context"
"net"
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
) )
const ( const (
@@ -30,12 +31,12 @@ type circuitBreakerDialer struct {
resetFailuresTimeout time.Duration resetFailuresTimeout time.Duration
} }
func (c *circuitBreakerDialer) Dial(network, address string) (net.Conn, error) { func (c *circuitBreakerDialer) Dial(network, address string) (essentials.Conn, error) {
return c.DialContext(context.Background(), network, address) return c.DialContext(context.Background(), network, address)
} }
func (c *circuitBreakerDialer) DialContext(ctx context.Context, func (c *circuitBreakerDialer) DialContext(ctx context.Context,
network, address string) (net.Conn, error) { network, address string) (essentials.Conn, error) {
switch atomic.LoadUint32(&c.state) { switch atomic.LoadUint32(&c.state) {
case circuitBreakerStateClosed: case circuitBreakerStateClosed:
return c.doClosed(ctx, network, address) return c.doClosed(ctx, network, address)
@@ -47,7 +48,7 @@ func (c *circuitBreakerDialer) DialContext(ctx context.Context,
} }
func (c *circuitBreakerDialer) doClosed(ctx context.Context, func (c *circuitBreakerDialer) doClosed(ctx context.Context,
network, address string) (net.Conn, error) { network, address string) (essentials.Conn, error) {
conn, err := c.Dialer.DialContext(ctx, network, address) conn, err := c.Dialer.DialContext(ctx, network, address)
select { select {
@@ -78,7 +79,8 @@ func (c *circuitBreakerDialer) doClosed(ctx context.Context,
return conn, err // nolint: wrapcheck return conn, err // nolint: wrapcheck
} }
func (c *circuitBreakerDialer) doHalfOpened(ctx context.Context, network, address string) (net.Conn, error) { func (c *circuitBreakerDialer) doHalfOpened(ctx context.Context,
network, address string) (essentials.Conn, error) {
if !atomic.CompareAndSwapUint32(&c.halfOpenAttempts, 0, 1) { if !atomic.CompareAndSwapUint32(&c.halfOpenAttempts, 0, 1) {
return nil, ErrCircuitBreakerOpened return nil, ErrCircuitBreakerOpened
} }
+2 -2
View File
@@ -21,7 +21,7 @@ type CircuitBreakerTestSuite struct {
mutex sync.Mutex mutex sync.Mutex
ctx context.Context ctx context.Context
ctxCancel context.CancelFunc ctxCancel context.CancelFunc
connMock *testlib.NetConnMock connMock *testlib.EssentialsConnMock
baseDialerMock *DialerMock baseDialerMock *DialerMock
} }
@@ -29,7 +29,7 @@ func (suite *CircuitBreakerTestSuite) SetupTest() {
suite.mutex = sync.Mutex{} suite.mutex = sync.Mutex{}
suite.ctx, suite.ctxCancel = context.WithCancel(context.Background()) suite.ctx, suite.ctxCancel = context.WithCancel(context.Background())
suite.baseDialerMock = &DialerMock{} suite.baseDialerMock = &DialerMock{}
suite.connMock = &testlib.NetConnMock{} suite.connMock = &testlib.EssentialsConnMock{}
suite.d = newCircuitBreakerDialer(suite.baseDialerMock, suite.d = newCircuitBreakerDialer(suite.baseDialerMock,
3, 100*time.Millisecond, 50*time.Millisecond) 3, 100*time.Millisecond, 50*time.Millisecond)
} }
+5 -3
View File
@@ -5,17 +5,19 @@ import (
"fmt" "fmt"
"net" "net"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
) )
type defaultDialer struct { type defaultDialer struct {
net.Dialer net.Dialer
} }
func (d *defaultDialer) Dial(network, address string) (net.Conn, error) { func (d *defaultDialer) Dial(network, address string) (essentials.Conn, error) {
return d.DialContext(context.Background(), network, address) return d.DialContext(context.Background(), network, address)
} }
func (d *defaultDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { func (d *defaultDialer) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) {
switch network { switch network {
case "tcp", "tcp4", "tcp6": // nolint: goconst case "tcp", "tcp4", "tcp6": // nolint: goconst
default: default:
@@ -34,7 +36,7 @@ func (d *defaultDialer) DialContext(ctx context.Context, network, address string
return nil, fmt.Errorf("cannot set socket options: %w", err) return nil, fmt.Errorf("cannot set socket options: %w", err)
} }
return conn, nil return conn.(essentials.Conn), nil
} }
// NewDefaultDialer build a new dialer which dials bypassing proxies // NewDefaultDialer build a new dialer which dials bypassing proxies
+4 -3
View File
@@ -20,8 +20,9 @@ package network
import ( import (
"context" "context"
"errors" "errors"
"net"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
) )
const ( const (
@@ -95,6 +96,6 @@ var (
// Dialer defines an interface which is required to bootstrap a network // Dialer defines an interface which is required to bootstrap a network
// instance from. // instance from.
type Dialer interface { type Dialer interface {
Dial(network, address string) (net.Conn, error) Dial(network, address string) (essentials.Conn, error)
DialContext(ctx context.Context, network, address string) (net.Conn, error) DialContext(ctx context.Context, network, address string) (essentials.Conn, error)
} }
+5 -5
View File
@@ -2,8 +2,8 @@ package network
import ( import (
"context" "context"
"net"
"github.com/9seconds/mtg/v2/essentials"
"github.com/stretchr/testify/mock" "github.com/stretchr/testify/mock"
) )
@@ -11,14 +11,14 @@ type DialerMock struct {
mock.Mock mock.Mock
} }
func (d *DialerMock) Dial(network, address string) (net.Conn, error) { func (d *DialerMock) Dial(network, address string) (essentials.Conn, error) {
args := d.Called(network, address) args := d.Called(network, address)
return args.Get(0).(net.Conn), args.Error(1) // nolint: wrapcheck return args.Get(0).(essentials.Conn), args.Error(1) // nolint: wrapcheck
} }
func (d *DialerMock) DialContext(ctx context.Context, network, address string) (net.Conn, error) { func (d *DialerMock) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) {
args := d.Called(ctx, network, address) args := d.Called(ctx, network, address)
return args.Get(0).(net.Conn), args.Error(1) // nolint: wrapcheck return args.Get(0).(essentials.Conn), args.Error(1) // nolint: wrapcheck
} }
+8 -5
View File
@@ -8,6 +8,7 @@ import (
"net/url" "net/url"
"strings" "strings"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/network" "github.com/9seconds/mtg/v2/network"
socks5 "github.com/armon/go-socks5" socks5 "github.com/armon/go-socks5"
"github.com/mccutchen/go-httpbin/httpbin" "github.com/mccutchen/go-httpbin/httpbin"
@@ -18,16 +19,16 @@ type DialerMock struct {
mock.Mock mock.Mock
} }
func (d *DialerMock) Dial(network, address string) (net.Conn, error) { func (d *DialerMock) Dial(network, address string) (essentials.Conn, error) {
args := d.Called(network, address) args := d.Called(network, address)
return args.Get(0).(net.Conn), args.Error(1) // nolint: wrapcheck return args.Get(0).(essentials.Conn), args.Error(1) // nolint: wrapcheck
} }
func (d *DialerMock) DialContext(ctx context.Context, network, address string) (net.Conn, error) { func (d *DialerMock) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) {
args := d.Called(ctx, network, address) args := d.Called(ctx, network, address)
return args.Get(0).(net.Conn), args.Error(1) // nolint: wrapcheck return args.Get(0).(essentials.Conn), args.Error(1) // nolint: wrapcheck
} }
type HTTPServerTestSuite struct { type HTTPServerTestSuite struct {
@@ -53,7 +54,9 @@ func (suite *HTTPServerTestSuite) MakeURL(path string) string {
func (suite *HTTPServerTestSuite) MakeHTTPClient(dialer network.Dialer) *http.Client { func (suite *HTTPServerTestSuite) MakeHTTPClient(dialer network.Dialer) *http.Client {
return &http.Client{ return &http.Client{
Transport: &http.Transport{ Transport: &http.Transport{
DialContext: dialer.DialContext, DialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
return dialer.DialContext(ctx, network, address) // nolint: wrapcheck
},
}, },
} }
} }
+4 -3
View File
@@ -4,19 +4,20 @@ import (
"context" "context"
"fmt" "fmt"
"math/rand" "math/rand"
"net"
"net/url" "net/url"
"github.com/9seconds/mtg/v2/essentials"
) )
type loadBalancedSocks5Dialer struct { type loadBalancedSocks5Dialer struct {
dialers []Dialer dialers []Dialer
} }
func (l loadBalancedSocks5Dialer) Dial(network, address string) (net.Conn, error) { func (l loadBalancedSocks5Dialer) Dial(network, address string) (essentials.Conn, error) {
return l.DialContext(context.Background(), network, address) return l.DialContext(context.Background(), network, address)
} }
func (l loadBalancedSocks5Dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { func (l loadBalancedSocks5Dialer) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) {
length := len(l.dialers) length := len(l.dialers)
start := rand.Intn(length) start := rand.Intn(length)
moved := false moved := false
+10 -6
View File
@@ -9,6 +9,7 @@ import (
"sync" "sync"
"time" "time"
"github.com/9seconds/mtg/v2/essentials"
"github.com/9seconds/mtg/v2/mtglib" "github.com/9seconds/mtg/v2/mtglib"
) )
@@ -30,11 +31,11 @@ type network struct {
dns *dnsResolver dns *dnsResolver
} }
func (n *network) Dial(protocol, address string) (net.Conn, error) { func (n *network) Dial(protocol, address string) (essentials.Conn, error) {
return n.DialContext(context.Background(), protocol, address) return n.DialContext(context.Background(), protocol, address)
} }
func (n *network) DialContext(ctx context.Context, protocol, address string) (net.Conn, error) { func (n *network) DialContext(ctx context.Context, protocol, address string) (essentials.Conn, error) {
host, port, _ := net.SplitHostPort(address) host, port, _ := net.SplitHostPort(address)
ips, err := n.dnsResolve(protocol, host) ips, err := n.dnsResolve(protocol, host)
@@ -46,7 +47,8 @@ func (n *network) DialContext(ctx context.Context, protocol, address string) (ne
ips[i], ips[j] = ips[j], ips[i] ips[i], ips[j] = ips[j], ips[i]
}) })
var conn net.Conn var conn essentials.Conn
for _, v := range ips { for _, v := range ips {
conn, err = n.dialer.DialContext(ctx, protocol, net.JoinHostPort(v, port)) conn, err = n.dialer.DialContext(ctx, protocol, net.JoinHostPort(v, port))
@@ -59,7 +61,7 @@ func (n *network) DialContext(ctx context.Context, protocol, address string) (ne
} }
func (n *network) MakeHTTPClient(dialFunc func(ctx context.Context, func (n *network) MakeHTTPClient(dialFunc func(ctx context.Context,
network, address string) (net.Conn, error)) *http.Client { network, address string) (essentials.Conn, error)) *http.Client {
if dialFunc == nil { if dialFunc == nil {
dialFunc = n.DialContext dialFunc = n.DialContext
} }
@@ -144,13 +146,15 @@ func NewNetwork(dialer Dialer,
func makeHTTPClient(userAgent string, func makeHTTPClient(userAgent string,
timeout time.Duration, timeout time.Duration,
dialFunc func(ctx context.Context, network, address string) (net.Conn, error)) *http.Client { dialFunc func(ctx context.Context, network, address string) (essentials.Conn, error)) *http.Client {
return &http.Client{ return &http.Client{
Timeout: timeout, Timeout: timeout,
Transport: networkHTTPTransport{ Transport: networkHTTPTransport{
userAgent: userAgent, userAgent: userAgent,
next: &http.Transport{ next: &http.Transport{
DialContext: dialFunc, DialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
return dialFunc(ctx, network, address)
},
}, },
}, },
} }
+146 -5
View File
@@ -1,21 +1,162 @@
package network package network
import ( import (
"context"
"fmt" "fmt"
"io"
"net"
"net/url" "net/url"
"golang.org/x/net/proxy" "github.com/9seconds/mtg/v2/essentials"
"github.com/txthinking/socks5"
) )
type socks5Dialer struct {
Dialer
username []byte
password []byte
proxyAddress string
}
func (s socks5Dialer) Dial(network, address string) (essentials.Conn, error) {
return s.DialContext(context.Background(), network, address)
}
func (s socks5Dialer) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) {
switch network {
case "tcp", "tcp4", "tcp6":
default:
return nil, fmt.Errorf("%s network type is not supported", network)
}
conn, err := s.Dialer.DialContext(ctx, network, s.proxyAddress)
if err != nil {
return nil, fmt.Errorf("cannot dial to the proxy: %w", err)
}
if err := s.handshake(conn); err != nil {
conn.Close()
return nil, fmt.Errorf("cannot perform a handshake: %w", err)
}
if err := s.connect(conn, address); err != nil {
conn.Close()
return nil, fmt.Errorf("cannot connect to a destination host %s: %w", address, err)
}
return conn, nil
}
func (s socks5Dialer) handshake(conn io.ReadWriter) error {
authMethod := socks5.MethodUsernamePassword
if len(s.username)+len(s.password) == 0 {
authMethod = socks5.MethodNone
}
if err := s.handshakeNegotiation(conn, authMethod); err != nil {
return fmt.Errorf("cannot perform negotiation: %w", err)
}
if authMethod == socks5.MethodNone {
return nil
}
if err := s.handshakeAuth(conn); err != nil {
return fmt.Errorf("cannot authenticate: %w", err)
}
return nil
}
func (s socks5Dialer) handshakeNegotiation(conn io.ReadWriter, authMethod byte) error {
request := socks5.NewNegotiationRequest([]byte{authMethod})
if _, err := request.WriteTo(conn); err != nil {
return fmt.Errorf("cannot send request: %w", err)
}
response, err := socks5.NewNegotiationReplyFrom(conn)
if err != nil {
return fmt.Errorf("cannot read response: %w", err)
}
if response.Method != authMethod {
return fmt.Errorf("%v is unsupported auth method", authMethod)
}
return nil
}
func (s socks5Dialer) handshakeAuth(conn io.ReadWriter) error {
request := socks5.NewUserPassNegotiationRequest(s.username, s.password)
if _, err := request.WriteTo(conn); err != nil {
return fmt.Errorf("cannot send a request: %w", err)
}
response, err := socks5.NewUserPassNegotiationReplyFrom(conn)
if err != nil {
return fmt.Errorf("cannot read a response: %w", err)
}
if response.Status != socks5.UserPassStatusSuccess {
return fmt.Errorf("authenticate has failed: %v", response.Status)
}
return nil
}
func (s socks5Dialer) connect(conn io.ReadWriter, address string) error {
addrType, host, port, err := socks5.ParseAddress(address)
if err != nil {
return fmt.Errorf("cannot parse address: %w", err)
}
if addrType == socks5.ATYPDomain {
host = host[1:]
}
request := socks5.NewRequest(socks5.CmdConnect, addrType, host, port)
if _, err := request.WriteTo(conn); err != nil {
return fmt.Errorf("cannot send a request: %w", err)
}
response, err := socks5.NewReplyFrom(conn)
if err != nil {
return fmt.Errorf("cannot read a response: %w", err)
}
if response.Rep != socks5.RepSuccess {
return fmt.Errorf("unsuccessful request: %v", response.Rep)
}
return nil
}
// NewSocks5Dialer build a new dialer from a given one (so, in theory // NewSocks5Dialer build a new dialer from a given one (so, in theory
// you can chain here). Proxy parameters are passed with URI in a form of: // you can chain here). Proxy parameters are passed with URI in a form of:
// //
// socks5://[user:[password]]@host:port // socks5://[user:[password]]@host:port
func NewSocks5Dialer(baseDialer Dialer, proxyURL *url.URL) (Dialer, error) { func NewSocks5Dialer(baseDialer Dialer, proxyURL *url.URL) (Dialer, error) {
rv, err := proxy.FromURL(proxyURL, baseDialer) if _, _, err := net.SplitHostPort(proxyURL.Host); err != nil {
if err != nil { return nil, fmt.Errorf("incorrect url %s", proxyURL.Redacted())
return nil, fmt.Errorf("cannot initialize socks5 proxy dialer: %w", err)
} }
return rv.(Dialer), nil dialer := socks5Dialer{
Dialer: baseDialer,
proxyAddress: proxyURL.Host,
}
if proxyURL.User != nil {
password, isSet := proxyURL.User.Password()
if isSet {
dialer.username = []byte(proxyURL.User.Username())
dialer.password = []byte(password)
}
}
return dialer, nil
} }
+1 -1
View File
@@ -55,7 +55,7 @@ func (suite *Socks5TestSuite) TestRequestOk() {
suite.Equal(http.StatusOK, resp.StatusCode) suite.Equal(http.StatusOK, resp.StatusCode)
} }
func TestSocks5TestSuite(t *testing.T) { func TestSocks5(t *testing.T) {
t.Parallel() t.Parallel()
suite.Run(t, &Socks5TestSuite{}) suite.Run(t, &Socks5TestSuite{})
} }