Add v2 network package

This commit is contained in:
9seconds
2026-02-27 15:50:21 +01:00
parent 282896be09
commit 42927c8bdc
15 changed files with 714 additions and 0 deletions
+25
View File
@@ -24,3 +24,28 @@ type Conn interface {
CloseableReader CloseableReader
CloseableWriter 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}
}
+1
View File
@@ -46,6 +46,7 @@ require (
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/client_model v0.6.2 // indirect
github.com/rogpeppe/go-internal v1.14.1 // 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 github.com/txthinking/runnergroup v0.0.0-20250224021307-5864ffeb65ae // indirect
go.yaml.in/yaml/v2 v2.4.3 // indirect go.yaml.in/yaml/v2 v2.4.3 // indirect
golang.org/x/sync v0.19.0 // indirect golang.org/x/sync v0.19.0 // indirect
+2
View File
@@ -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.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 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= 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-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 h1:ArVM1jICfm7g4E4dBet+KHUFMLuxmj1Nxdp/tr3ByCU=
github.com/txthinking/runnergroup v0.0.0-20250224021307-5864ffeb65ae/go.mod h1:cldYm15/XHcGt7ndItnEWHwFZo7dinU+2QoyjfErhsI= github.com/txthinking/runnergroup v0.0.0-20250224021307-5864ffeb65ae/go.mod h1:cldYm15/XHcGt7ndItnEWHwFZo7dinU+2QoyjfErhsI=
+49
View File
@@ -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{})
}
+83
View File
@@ -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{})
}
+106
View File
@@ -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()
}
+14
View File
@@ -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
}
+45
View File
@@ -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
}
+70
View File
@@ -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
}
+88
View File
@@ -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,
}
}
+36
View File
@@ -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
}
+27
View File
@@ -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
}
+31
View File
@@ -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
}
+10
View File
@@ -0,0 +1,10 @@
//go:build windows
// +build windows
package network
import "syscall"
func setSocketReuseAddrPort(conn syscall.RawConn) error {
return nil
}
+127
View File
@@ -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{})
}