mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 08:54:02 +03:00
Add v2 network package
This commit is contained in:
@@ -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}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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{})
|
||||
}
|
||||
@@ -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{})
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
//go:build windows
|
||||
// +build windows
|
||||
|
||||
package network
|
||||
|
||||
import "syscall"
|
||||
|
||||
func setSocketReuseAddrPort(conn syscall.RawConn) error {
|
||||
return nil
|
||||
}
|
||||
@@ -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{})
|
||||
}
|
||||
Reference in New Issue
Block a user