diff --git a/essentials/conns.go b/essentials/conns.go index 5b42efe..1917b6b 100644 --- a/essentials/conns.go +++ b/essentials/conns.go @@ -24,3 +24,28 @@ type Conn interface { CloseableReader CloseableWriter } + +type netConnWrapper struct { + net.Conn +} + +func (n netConnWrapper) CloseRead() error { + if conn, ok := n.Conn.(CloseableReader); ok { + return conn.CloseRead() + } + + return n.Close() +} + +func (n netConnWrapper) CloseWrite() error { + if conn, ok := n.Conn.(CloseableWriter); ok { + return conn.CloseWrite() + } + + return n.Close() +} + +// WrapConn wraps a generic [net.Conn] into Conn. +func WrapNetConn(conn net.Conn) Conn { + return netConnWrapper{conn} +} diff --git a/go.mod b/go.mod index b78ee66..96d1216 100644 --- a/go.mod +++ b/go.mod @@ -46,6 +46,7 @@ require ( github.com/pmezard/go-difflib v1.0.0 // indirect github.com/prometheus/client_model v0.6.2 // indirect github.com/rogpeppe/go-internal v1.14.1 // indirect + github.com/things-go/go-socks5 v0.1.0 // indirect github.com/txthinking/runnergroup v0.0.0-20250224021307-5864ffeb65ae // indirect go.yaml.in/yaml/v2 v2.4.3 // indirect golang.org/x/sync v0.19.0 // indirect diff --git a/go.sum b/go.sum index 6c41391..a9cdd4a 100644 --- a/go.sum +++ b/go.sum @@ -89,6 +89,8 @@ github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXl github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/things-go/go-socks5 v0.1.0 h1:4f5dz0iMQ6cA4wseFmyLmCHmg3SWJTW92ndrKS6oERg= +github.com/things-go/go-socks5 v0.1.0/go.mod h1:Riabiyu52kLsla0YmJqunt1c1JEl6iXSr4bRd7swFEA= github.com/txthinking/runnergroup v0.0.0-20210608031112-152c7c4432bf/go.mod h1:CLUSJbazqETbaR+i0YAhXBICV9TrKH93pziccMhmhpM= github.com/txthinking/runnergroup v0.0.0-20250224021307-5864ffeb65ae h1:ArVM1jICfm7g4E4dBet+KHUFMLuxmj1Nxdp/tr3ByCU= github.com/txthinking/runnergroup v0.0.0-20250224021307-5864ffeb65ae/go.mod h1:cldYm15/XHcGt7ndItnEWHwFZo7dinU+2QoyjfErhsI= diff --git a/network/v2/base_http_test.go b/network/v2/base_http_test.go new file mode 100644 index 0000000..1d322bb --- /dev/null +++ b/network/v2/base_http_test.go @@ -0,0 +1,49 @@ +package network_test + +import ( + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/9seconds/mtg/v2/network/v2" + "github.com/stretchr/testify/suite" +) + +type BaseHTTPTestSuite struct { + suite.Suite + + http *httptest.Server + client *http.Client +} + +func (suite *BaseHTTPTestSuite) SetupSuite() { + suite.http = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + w.Write([]byte(r.Header.Get("User-Agent"))) + })) +} + +func (suite *BaseHTTPTestSuite) SetupTest() { + suite.client = network.New(nil, "mtg/1", 0, 0, 0).MakeHTTPClient(nil) +} + +func (suite *BaseHTTPTestSuite) TestGet() { + resp, err := suite.client.Get(suite.http.URL) + suite.NoError(err) + + defer resp.Body.Close() + + data, err := io.ReadAll(resp.Body) + suite.NoError(err) + suite.Equal("mtg/1", string(data)) +} + +func (suite *BaseHTTPTestSuite) TearDownSuite() { + suite.http.Close() +} + +func TestBaseHTTP(t *testing.T) { + t.Parallel() + suite.Run(t, &BaseHTTPTestSuite{}) +} diff --git a/network/v2/base_network_test.go b/network/v2/base_network_test.go new file mode 100644 index 0000000..f9729e4 --- /dev/null +++ b/network/v2/base_network_test.go @@ -0,0 +1,83 @@ +package network_test + +import ( + "context" + "testing" + + "github.com/9seconds/mtg/v2/network/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/suite" +) + +type BaseNetworkTestSuite struct { + EchoServerTestSuite + + net network.Network +} + +func (suite *BaseNetworkTestSuite) SetupSuite() { + suite.EchoServerTestSuite.SetupSuite() + + suite.net = network.New(nil, "agent", 0, 0, 0) +} + +func (suite *BaseNetworkTestSuite) TestDialUnknownNetwork() { + testData := []string{ + "udp", + "udp4", + "udp6", + "unix", + } + + for _, name := range testData { + suite.T().Run(name, func(t *testing.T) { + _, err := suite.net.Dial(name, suite.EchoServerAddr()) + assert.Error(t, err) + }) + } +} + +func (suite *BaseNetworkTestSuite) TestDial() { + conn, err := suite.net.Dial("tcp4", suite.EchoServerAddr()) + suite.NoError(err) + + buf := []byte{1, 2, 3, 4, 5} + n, err := conn.Write(buf) + suite.Equal(5, n) + suite.NoError(err) + + another := make([]byte, len(buf)) + n, err = conn.Read(another) + suite.NoError(err) + suite.Equal(len(another), n) + suite.Equal(buf, another) +} + +func (suite *BaseNetworkTestSuite) TestDialContextOk() { + conn, err := suite.net.DialContext(context.Background(), "tcp4", suite.EchoServerAddr()) + suite.NoError(err) + + buf := []byte{1, 2, 3, 4, 5} + n, err := conn.Write(buf) + suite.Equal(5, n) + suite.NoError(err) + + another := make([]byte, len(buf)) + n, err = conn.Read(another) + suite.NoError(err) + suite.Equal(len(another), n) + suite.Equal(buf, another) +} + +func (suite *BaseNetworkTestSuite) TestDialContextClosed() { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := suite.net.DialContext(ctx, "tcp4", suite.EchoServerAddr()) + suite.ErrorIs(err, ctx.Err()) +} + +func TestNetworkBase(t *testing.T) { + t.Parallel() + suite.Run(t, &BaseNetworkTestSuite{}) +} diff --git a/network/v2/echo_server_test.go b/network/v2/echo_server_test.go new file mode 100644 index 0000000..cfcead8 --- /dev/null +++ b/network/v2/echo_server_test.go @@ -0,0 +1,106 @@ +package network_test + +import ( + "context" + "io" + "net" + "sync" + + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" +) + +type EchoServer struct { + wg sync.WaitGroup + ctx context.Context + ctxCancel context.CancelFunc + listener net.Listener +} + +func (e *EchoServer) Run() { + e.wg.Go(func() { + <-e.ctx.Done() + e.listener.Close() + }) + + e.wg.Go(func() { + for { + conn, err := e.listener.Accept() + if err != nil { + return + } + + e.wg.Go(func() { + <-e.ctx.Done() + conn.Close() + }) + e.wg.Go(func() { + e.process(conn) + }) + } + }) +} + +func (e *EchoServer) Stop() { + e.ctxCancel() + e.wg.Wait() +} + +func (e *EchoServer) Addr() string { + return e.listener.Addr().String() +} + +func (e *EchoServer) process(conn io.ReadWriter) { + buf := [4096]byte{} + + for { + select { + case <-e.ctx.Done(): + return + default: + } + + n, err := conn.Read(buf[:]) + if err != nil { + return + } + + select { + case <-e.ctx.Done(): + return + default: + } + + if _, err = conn.Write(buf[:n]); err != nil { + return + } + } +} + +type EchoServerTestSuite struct { + suite.Suite + + echoServer *EchoServer +} + +func (suite *EchoServerTestSuite) SetupSuite() { + ctx, cancel := context.WithCancel(context.Background()) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(suite.T(), err) + + suite.echoServer = &EchoServer{ + ctx: ctx, + ctxCancel: cancel, + listener: listener, + } + suite.echoServer.Run() +} + +func (suite *EchoServerTestSuite) TearDownSuite() { + suite.echoServer.Stop() +} + +func (suite *EchoServerTestSuite) EchoServerAddr() string { + return suite.echoServer.Addr() +} diff --git a/network/v2/http.go b/network/v2/http.go new file mode 100644 index 0000000..58e6a5e --- /dev/null +++ b/network/v2/http.go @@ -0,0 +1,14 @@ +package network + +import "net/http" + +type networkHTTPTransport struct { + userAgent string + next http.RoundTripper +} + +func (n networkHTTPTransport) RoundTrip(req *http.Request) (*http.Response, error) { + req.Header.Set("User-Agent", n.userAgent) + + return n.next.RoundTrip(req) //nolint: wrapcheck +} diff --git a/network/v2/init.go b/network/v2/init.go new file mode 100644 index 0000000..7cc4302 --- /dev/null +++ b/network/v2/init.go @@ -0,0 +1,45 @@ +// Network contains a default implementation of the network. +// +// Please see [mtglib.Network] interface to get some basic idea behind this +// abstraction. +// +// This implementation is more simple that v1 because life shows that all +// this complexity, especially around circuit breakers and DoH is not really +// required. There is no chance that if DNS address is spoofed, that real +// IP would work as expected. +package network + +import ( + "errors" + "net" + "time" + + "github.com/9seconds/mtg/v2/mtglib" +) + +const ( + // DefaultTimeout is a default timeout for establishing TCP connection. + DefaultTimeout = 10 * time.Second + + // DefaultHTTPTimeout defines a default timeout for making HTTP request. + DefaultHTTPTimeout = 10 * time.Second + + // DefaultIdleTimeout defines a timeout for idle HTTP connections + DefaultIdleTimeout = time.Minute + + // DefaultTCPKeepAlivePeriod defines a time period between 2 consecuitive + // probes. + DefaultTCPKeepAlivePeriod = 10 * time.Second + + // tcpLingerTimeout defines a number of seconds to wait for sending + // unacknowledged data. + tcpLingerTimeout = 1 +) + +var ErrCannotDial = errors.New("cannot dial to any address") + +type Network interface { + mtglib.Network + + NativeDialer() *net.Dialer +} diff --git a/network/v2/multi_network.go b/network/v2/multi_network.go new file mode 100644 index 0000000..45e501c --- /dev/null +++ b/network/v2/multi_network.go @@ -0,0 +1,70 @@ +package network + +import ( + "context" + "errors" + "math/rand" + "net" + "net/http" + + "github.com/9seconds/mtg/v2/essentials" +) + +type multiNetwork struct { + networks []Network +} + +func (m multiNetwork) Dial(network, address string) (essentials.Conn, error) { + return m.DialContext(context.Background(), network, address) +} + +func (m multiNetwork) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) { + networks := m.networks + + if len(networks) > 1 { + networks = make([]Network, len(m.networks)) + copy(networks, m.networks) + + rand.Shuffle(len(m.networks), func(i, j int) { + networks[i], networks[j] = networks[j], networks[i] + }) + } + + errs := make([]error, 1, len(networks)+1) + errs[0] = ErrCannotDial + + for _, ntw := range networks { + conn, err := ntw.DialContext(ctx, network, address) + if err == nil { + return conn, nil + } + + errs = append(errs, err) + } + + return nil, errors.Join(errs...) +} + +func (m multiNetwork) NativeDialer() *net.Dialer { + return m.networks[0].NativeDialer() +} + +func (m multiNetwork) MakeHTTPClient( + dialFunc func(context.Context, string, string) (essentials.Conn, error), +) *http.Client { + if dialFunc == nil { + dialFunc = m.DialContext + } + + return m.networks[0].MakeHTTPClient(dialFunc) +} + +func Join(networks ...Network) (Network, error) { + if len(networks) == 0 { + return nil, errors.New("cannot join no networks") + } + + return multiNetwork{ + networks: networks, + }, nil +} diff --git a/network/v2/network.go b/network/v2/network.go new file mode 100644 index 0000000..9921f8b --- /dev/null +++ b/network/v2/network.go @@ -0,0 +1,88 @@ +package network + +import ( + "context" + "fmt" + "net" + "net/http" + "time" + + "github.com/9seconds/mtg/v2/essentials" +) + +type network struct { + net.Dialer + + httpTimeout time.Duration + idleTimeout time.Duration + userAgent string +} + +func (n *network) Dial(network, address string) (essentials.Conn, error) { + return n.DialContext(context.Background(), network, address) +} + +func (n *network) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) { + switch network { + case "tcp", "tcp4", "tcp6": + default: + return nil, fmt.Errorf("unsupported network %s", network) + } + + conn, err := n.Dialer.DialContext(ctx, network, address) + if err != nil { + return nil, err + } + + tcpConn := conn.(*net.TCPConn) + + return tcpConn, setCommonSocketOptions(tcpConn) +} + +func (n *network) MakeHTTPClient( + dialFunc func(context.Context, string, string) (essentials.Conn, error), +) *http.Client { + if dialFunc == nil { + dialFunc = n.DialContext + } + + return &http.Client{ + Timeout: n.httpTimeout, + Transport: networkHTTPTransport{ + userAgent: n.userAgent, + next: &http.Transport{ + IdleConnTimeout: n.idleTimeout, + DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + return dialFunc(ctx, network, address) + }, + }, + }, + } +} + +func (n *network) NativeDialer() *net.Dialer { + return &n.Dialer +} + +func New( + dnsResolver *net.Resolver, + userAgent string, + tcpTimeout, + httpTimeout, + idleTimeout time.Duration, +) Network { + if dnsResolver == nil { + dnsResolver = net.DefaultResolver + } + + return &network{ + Dialer: net.Dialer{ + Timeout: tcpTimeout, + Resolver: dnsResolver, + FallbackDelay: -1, + }, + userAgent: userAgent, + idleTimeout: idleTimeout, + httpTimeout: httpTimeout, + } +} diff --git a/network/v2/proxy_network.go b/network/v2/proxy_network.go new file mode 100644 index 0000000..c27162d --- /dev/null +++ b/network/v2/proxy_network.go @@ -0,0 +1,36 @@ +package network + +import ( + "context" + "fmt" + "net/url" + + "github.com/9seconds/mtg/v2/essentials" + "golang.org/x/net/proxy" +) + +type proxyNetwork struct { + Network + client proxy.ContextDialer +} + +func (p proxyNetwork) DialContext(ctx context.Context, network, address string) (essentials.Conn, error) { + conn, err := p.client.DialContext(ctx, network, address) + if err != nil { + return nil, err + } + + return essentials.WrapNetConn(conn), nil +} + +func NewProxyNetwork(base Network, proxyURL *url.URL) (*proxyNetwork, error) { + socks, err := proxy.FromURL(proxyURL, base.NativeDialer()) + if err != nil { + return nil, fmt.Errorf("cannot build proxy dialer: %w", err) + } + + return &proxyNetwork{ + Network: base, + client: socks.(proxy.ContextDialer), + }, nil +} diff --git a/network/v2/sockopts.go b/network/v2/sockopts.go new file mode 100644 index 0000000..378db00 --- /dev/null +++ b/network/v2/sockopts.go @@ -0,0 +1,27 @@ +package network + +import ( + "fmt" + "net" +) + +func setCommonSocketOptions(conn *net.TCPConn) error { + if err := conn.SetKeepAlivePeriod(DefaultTCPKeepAlivePeriod); err != nil { + return fmt.Errorf("cannot set time period of TCP keepalive probes: %w", err) + } + + if err := conn.SetLinger(tcpLingerTimeout); err != nil { + return fmt.Errorf("cannot set TCP linger timeout: %w", err) + } + + rawConn, err := conn.SyscallConn() + if err != nil { + return fmt.Errorf("cannot get underlying raw connection: %w", err) + } + + if err := setSocketReuseAddrPort(rawConn); err != nil { + return fmt.Errorf("cannot setup SO_REUSEADDR/PORT: %w", err) + } + + return nil +} diff --git a/network/v2/sockopts_unix.go b/network/v2/sockopts_unix.go new file mode 100644 index 0000000..75df3af --- /dev/null +++ b/network/v2/sockopts_unix.go @@ -0,0 +1,31 @@ +//go:build !windows +// +build !windows + +package network + +import ( + "fmt" + "syscall" + + "golang.org/x/sys/unix" +) + +func setSocketReuseAddrPort(conn syscall.RawConn) error { + var err error + + conn.Control(func(fd uintptr) { //nolint: errcheck + err = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEADDR, 1) + if err != nil { + err = fmt.Errorf("cannot set SO_REUSEADDR: %w", err) + + return + } + + err = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1) + if err != nil { + err = fmt.Errorf("cannot set SO_REUSEPORT: %w", err) + } + }) + + return err +} diff --git a/network/v2/sockopts_windows.go b/network/v2/sockopts_windows.go new file mode 100644 index 0000000..32a702a --- /dev/null +++ b/network/v2/sockopts_windows.go @@ -0,0 +1,10 @@ +//go:build windows +// +build windows + +package network + +import "syscall" + +func setSocketReuseAddrPort(conn syscall.RawConn) error { + return nil +} diff --git a/network/v2/socks_proxy_test.go b/network/v2/socks_proxy_test.go new file mode 100644 index 0000000..55262ac --- /dev/null +++ b/network/v2/socks_proxy_test.go @@ -0,0 +1,127 @@ +package network_test + +import ( + "net" + "net/url" + "sync" + "testing" + + "github.com/9seconds/mtg/v2/network/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/things-go/go-socks5" +) + +type SocksProxyTestSuite struct { + EchoServerTestSuite + + wg sync.WaitGroup + baseNetwork network.Network + + noAuthURL *url.URL + authURL *url.URL + + noAuthListener net.Listener + authListener net.Listener + + noAuthServer *socks5.Server + authServer *socks5.Server +} + +func (suite *SocksProxyTestSuite) SetupSuite() { + suite.EchoServerTestSuite.SetupSuite() + + listener, err := net.Listen("tcp4", "127.0.0.1:0") + require.NoError(suite.T(), err) + suite.noAuthListener = listener + + listener, err = net.Listen("tcp4", "127.0.0.1:0") + require.NoError(suite.T(), err) + suite.authListener = listener + + suite.noAuthServer = socks5.NewServer() + suite.wg.Go(func() { + suite.noAuthServer.Serve(suite.noAuthListener) + }) + + suite.authServer = socks5.NewServer( + socks5.WithAuthMethods([]socks5.Authenticator{ + socks5.UserPassAuthenticator{ + Credentials: socks5.StaticCredentials{ + "user": "pass", + }, + }, + })) + suite.wg.Go(func() { + suite.authServer.Serve(suite.authListener) + }) + + parsed, err := url.Parse("socks5://" + suite.noAuthListener.Addr().String()) + require.NoError(suite.T(), err) + suite.noAuthURL = parsed + + parsed, err = url.Parse("socks5://user:pass@" + suite.authListener.Addr().String()) + require.NoError(suite.T(), err) + suite.authURL = parsed + + suite.baseNetwork = network.New(nil, "mtg", 0, 0, 0) +} + +func (suite *SocksProxyTestSuite) TestIncorrectSchema() { + parsed, err := url.Parse("http://hello") + suite.NoError(err) + + _, err = network.NewProxyNetwork(suite.baseNetwork, parsed) + suite.Error(err) +} + +func (suite *SocksProxyTestSuite) TestRead() { + testData := map[string][]*url.URL{ + "noAuth": {suite.noAuthURL}, + "auth": {suite.authURL}, + "both": {suite.noAuthURL, suite.authURL}, + } + + for name, proxies := range testData { + suite.T().Run(name, func(t *testing.T) { + proxyNetworks := []network.Network{} + + for _, u := range proxies { + value, err := network.NewProxyNetwork(suite.baseNetwork, u) + assert.NoError(t, err) + proxyNetworks = append(proxyNetworks, value) + } + + netw, err := network.Join(proxyNetworks...) + assert.NoError(t, err) + + conn, err := netw.Dial("tcp4", suite.EchoServerAddr()) + assert.NoError(t, err) + + data := []byte{1, 2, 3} + n, err := conn.Write(data) + assert.NoError(t, err) + assert.Equal(t, len(data), n) + + toRead := []byte{1, 2, 3, 4, 5} + n, err = conn.Read(toRead) + assert.NoError(t, err) + assert.Equal(t, len(data), n) + assert.Equal(t, data, toRead[:n]) + assert.NotEqual(t, data, toRead) + }) + } +} + +func (suite *SocksProxyTestSuite) TearDownSuite() { + suite.noAuthListener.Close() + suite.authListener.Close() + suite.wg.Wait() + suite.EchoServerTestSuite.TearDownSuite() +} + +func TestSocksProxy(t *testing.T) { + t.Parallel() + suite.Run(t, &SocksProxyTestSuite{}) +}