mirror of
https://github.com/qdm12/gluetun.git
synced 2026-07-22 10:26:26 +02:00
review tests and fix AI found issues
This commit is contained in:
@@ -91,14 +91,14 @@ func encodeBindData(addrType addrType, address string, port uint16) (
|
|||||||
return data, nil
|
return data, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func bindDataLength(addrType addrType, address string) (maxLength int) {
|
func bindDataLength(addrType addrType, address string) (maxLength uint) {
|
||||||
maxLength++ // address type
|
maxLength++ // address type
|
||||||
switch addrType {
|
switch addrType {
|
||||||
case ipv4:
|
case ipv4:
|
||||||
maxLength += net.IPv4len
|
maxLength += net.IPv4len
|
||||||
case domainName:
|
case domainName:
|
||||||
maxLength++ // domain name length
|
maxLength++ // domain name length
|
||||||
maxLength += len([]byte(address))
|
maxLength += uint(len([]byte(address)))
|
||||||
case ipv6:
|
case ipv6:
|
||||||
maxLength += net.IPv6len
|
maxLength += net.IPv6len
|
||||||
default:
|
default:
|
||||||
|
|||||||
+17
-12
@@ -8,7 +8,7 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Server struct {
|
type server struct {
|
||||||
username string
|
username string
|
||||||
password string
|
password string
|
||||||
address string
|
address string
|
||||||
@@ -23,8 +23,8 @@ type Server struct {
|
|||||||
stopping atomic.Bool
|
stopping atomic.Bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func newServer(settings Settings) *Server {
|
func newServer(settings Settings) *server {
|
||||||
return &Server{
|
return &server{
|
||||||
username: settings.Username,
|
username: settings.Username,
|
||||||
password: settings.Password,
|
password: settings.Password,
|
||||||
address: settings.Address,
|
address: settings.Address,
|
||||||
@@ -32,7 +32,7 @@ func newServer(settings Settings) *Server {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) Start(ctx context.Context) (runErr <-chan error, err error) {
|
func (s *server) Start(ctx context.Context) (runErr <-chan error, err error) {
|
||||||
s.socksConnCtx, s.socksConnCancel = context.WithCancel(context.Background())
|
s.socksConnCtx, s.socksConnCancel = context.WithCancel(context.Background())
|
||||||
config := &net.ListenConfig{}
|
config := &net.ListenConfig{}
|
||||||
s.listener, err = config.Listen(s.socksConnCtx, "tcp", s.address)
|
s.listener, err = config.Listen(s.socksConnCtx, "tcp", s.address)
|
||||||
@@ -56,7 +56,7 @@ func (s *Server) Start(ctx context.Context) (runErr <-chan error, err error) {
|
|||||||
return runErr, nil
|
return runErr, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) runServer(ready chan<- struct{},
|
func (s *server) runServer(ready chan<- struct{},
|
||||||
runErrCh chan<- error, done chan<- struct{},
|
runErrCh chan<- error, done chan<- struct{},
|
||||||
) {
|
) {
|
||||||
close(ready)
|
close(ready)
|
||||||
@@ -69,7 +69,7 @@ func (s *Server) runServer(ready chan<- struct{},
|
|||||||
connection, err := s.listener.Accept()
|
connection, err := s.listener.Accept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if !s.stopping.Load() {
|
if !s.stopping.Load() {
|
||||||
_ = s.Stop()
|
_ = s.stop()
|
||||||
runErrCh <- fmt.Errorf("accepting connection: %w", err)
|
runErrCh <- fmt.Errorf("accepting connection: %w", err)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
@@ -94,17 +94,22 @@ func (s *Server) runServer(ready chan<- struct{},
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) Stop() (err error) {
|
func (s *server) Stop() (err error) {
|
||||||
s.stopping.Store(true)
|
s.stopping.Store(true)
|
||||||
s.listening.Store(false)
|
err = s.stop()
|
||||||
err = s.listener.Close()
|
<-s.done // wait for run goroutine to finish
|
||||||
s.socksConnCancel() // stop ongoing socks connections
|
|
||||||
<-s.done // wait for run goroutine to finish
|
|
||||||
s.stopping.Store(false)
|
s.stopping.Store(false)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) listeningAddress() net.Addr {
|
func (s *server) stop() error {
|
||||||
|
s.listening.Store(false)
|
||||||
|
err := s.listener.Close()
|
||||||
|
s.socksConnCancel() // stop ongoing socks connections
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *server) listeningAddress() net.Addr {
|
||||||
if s.listening.Load() {
|
if s.listening.Load() {
|
||||||
return s.listener.Addr()
|
return s.listener.Addr()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -132,7 +132,8 @@ func (c *socksConn) handleRequest(ctx context.Context) error {
|
|||||||
return fmt.Errorf("writing successful %s response: %w", request.command, err)
|
return fmt.Errorf("writing successful %s response: %w", request.command, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
errc := make(chan error)
|
const capacity = 2 // if one goroutine fails, we don't want to leak the other one
|
||||||
|
errc := make(chan error, capacity)
|
||||||
go func() {
|
go func() {
|
||||||
_, err := io.Copy(c.clientConn, destinationConn)
|
_, err := io.Copy(c.clientConn, destinationConn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+89
-145
@@ -2,7 +2,6 @@ package socks5
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
@@ -17,8 +16,8 @@ import (
|
|||||||
|
|
||||||
type noopLogger struct{}
|
type noopLogger struct{}
|
||||||
|
|
||||||
func (noopLogger) Infof(string, ...interface{}) {}
|
func (noopLogger) Infof(string, ...any) {}
|
||||||
func (noopLogger) Warnf(string, ...interface{}) {}
|
func (noopLogger) Warnf(string, ...any) {}
|
||||||
|
|
||||||
func TestServerProxy(t *testing.T) {
|
func TestServerProxy(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
@@ -38,10 +37,10 @@ func TestServerProxy(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
// Backend TCP server: accepts one connection for the proxy to forward to.
|
// Backend TCP server: accepts one connection for the proxy to forward to.
|
||||||
backendListener, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", "127.0.0.1:0")
|
backendListener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", "127.0.0.1:0")
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
backendConnCh := make(chan net.Conn, 1)
|
backendConnCh := make(chan net.Conn)
|
||||||
go func() {
|
go func() {
|
||||||
conn, err := backendListener.Accept()
|
conn, err := backendListener.Accept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -56,7 +55,7 @@ func TestServerProxy(t *testing.T) {
|
|||||||
Address: "127.0.0.1:0",
|
Address: "127.0.0.1:0",
|
||||||
Logger: noopLogger{},
|
Logger: noopLogger{},
|
||||||
})
|
})
|
||||||
_, err = server.Start(context.Background())
|
_, err = server.Start(t.Context())
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
_ = server.Stop()
|
_ = server.Stop()
|
||||||
@@ -107,7 +106,7 @@ func dialSOCKS5(t *testing.T, proxyAddr, targetAddr, username, password string)
|
|||||||
targetPort, err := strconv.Atoi(portStr)
|
targetPort, err := strconv.Atoi(portStr)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
conn, err := (&net.Dialer{}).DialContext(context.Background(), "tcp", proxyAddr)
|
conn, err := (&net.Dialer{}).DialContext(t.Context(), "tcp", proxyAddr)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
var method authMethod
|
var method authMethod
|
||||||
@@ -180,11 +179,11 @@ func dialSOCKS5(t *testing.T, proxyAddr, targetAddr, username, password string)
|
|||||||
return conn
|
return conn
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNew(t *testing.T) {
|
func Test_newServer(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
settings Settings
|
settings Settings
|
||||||
expected *Server
|
expected *server
|
||||||
}{
|
}{
|
||||||
"with_auth": {
|
"with_auth": {
|
||||||
settings: Settings{
|
settings: Settings{
|
||||||
@@ -193,7 +192,7 @@ func TestNew(t *testing.T) {
|
|||||||
Address: "127.0.0.1:1080",
|
Address: "127.0.0.1:1080",
|
||||||
Logger: nil,
|
Logger: nil,
|
||||||
},
|
},
|
||||||
expected: &Server{
|
expected: &server{
|
||||||
username: "user",
|
username: "user",
|
||||||
password: "pass",
|
password: "pass",
|
||||||
address: "127.0.0.1:1080",
|
address: "127.0.0.1:1080",
|
||||||
@@ -205,7 +204,7 @@ func TestNew(t *testing.T) {
|
|||||||
Address: "127.0.0.1:1080",
|
Address: "127.0.0.1:1080",
|
||||||
Logger: nil,
|
Logger: nil,
|
||||||
},
|
},
|
||||||
expected: &Server{
|
expected: &server{
|
||||||
address: "127.0.0.1:1080",
|
address: "127.0.0.1:1080",
|
||||||
logger: nil,
|
logger: nil,
|
||||||
},
|
},
|
||||||
@@ -224,7 +223,7 @@ func TestNew(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestStartStop(t *testing.T) {
|
func Test_Server_StartStop(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
ctrl := gomock.NewController(t)
|
ctrl := gomock.NewController(t)
|
||||||
logger := NewMockLogger(ctrl)
|
logger := NewMockLogger(ctrl)
|
||||||
@@ -236,7 +235,7 @@ func TestStartStop(t *testing.T) {
|
|||||||
Logger: logger,
|
Logger: logger,
|
||||||
})
|
})
|
||||||
|
|
||||||
runErr, startErr := server.Start(context.Background())
|
runErr, startErr := server.Start(t.Context())
|
||||||
require.NoError(t, startErr)
|
require.NoError(t, startErr)
|
||||||
|
|
||||||
select {
|
select {
|
||||||
@@ -252,7 +251,7 @@ func TestStartStop(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestEncodeBindData(t *testing.T) {
|
func Test_encodeBindData(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
addrType addrType
|
addrType addrType
|
||||||
@@ -319,61 +318,47 @@ func TestEncodeBindData(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDecodeRequest(t *testing.T) {
|
func Test_decodeRequest(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
buildPacket func() []byte
|
packet []byte
|
||||||
expectedErr string
|
expectedErr string
|
||||||
validate func(*testing.T, request)
|
validate func(*testing.T, request)
|
||||||
}{
|
}{
|
||||||
"ipv4_valid": {
|
"ipv4_valid": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{socks5Version, byte(connect), 0, byte(ipv4), 127, 0, 0, 1, byte(0x1f), byte(0x90)},
|
||||||
packet := []byte{socks5Version, byte(connect), 0, byte(ipv4)}
|
validate: func(t *testing.T, request request) {
|
||||||
packet = append(packet, 127, 0, 0, 1)
|
|
||||||
packet = append(packet, byte(0x1f), byte(0x90))
|
|
||||||
return packet
|
|
||||||
},
|
|
||||||
validate: func(t *testing.T, req request) {
|
|
||||||
t.Helper()
|
t.Helper()
|
||||||
assert.Equal(t, connect, req.command)
|
assert.Equal(t, connect, request.command)
|
||||||
assert.Equal(t, "127.0.0.1", req.destination)
|
assert.Equal(t, "127.0.0.1", request.destination)
|
||||||
assert.Equal(t, uint16(8080), req.port)
|
assert.Equal(t, uint16(8080), request.port)
|
||||||
assert.Equal(t, ipv4, req.addressType)
|
assert.Equal(t, ipv4, request.addressType)
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
"domain_name_valid": {
|
"domain_name_valid": {
|
||||||
buildPacket: func() []byte {
|
packet: concatBytes(
|
||||||
packet := []byte{socks5Version, byte(connect), 0, byte(domainName)}
|
[]byte{socks5Version, byte(connect), 0, byte(domainName)},
|
||||||
domain := "example.com"
|
[]byte{byte(len("example.com"))},
|
||||||
packet = append(packet, byte(len(domain)))
|
[]byte("example.com"),
|
||||||
packet = append(packet, []byte(domain)...)
|
[]byte{0x00, 0x50},
|
||||||
packet = append(packet, byte(0x00), byte(0x50))
|
),
|
||||||
return packet
|
validate: func(t *testing.T, request request) {
|
||||||
},
|
|
||||||
validate: func(t *testing.T, req request) {
|
|
||||||
t.Helper()
|
t.Helper()
|
||||||
assert.Equal(t, "example.com", req.destination)
|
assert.Equal(t, "example.com", request.destination)
|
||||||
assert.Equal(t, uint16(80), req.port)
|
assert.Equal(t, uint16(80), request.port)
|
||||||
assert.Equal(t, domainName, req.addressType)
|
assert.Equal(t, domainName, request.addressType)
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
"version_mismatch": {
|
"version_mismatch": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{4, byte(connect), 0, byte(ipv4), 127, 0, 0, 1, 0, 0},
|
||||||
return []byte{4, byte(connect), 0, byte(ipv4), 127, 0, 0, 1, 0, 0}
|
|
||||||
},
|
|
||||||
expectedErr: "version is not supported",
|
expectedErr: "version is not supported",
|
||||||
},
|
},
|
||||||
"truncated_header": {
|
"truncated_header": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{socks5Version, byte(connect)},
|
||||||
return []byte{socks5Version, byte(connect)}
|
|
||||||
},
|
|
||||||
expectedErr: "reading header",
|
expectedErr: "reading header",
|
||||||
},
|
},
|
||||||
"unsupported_address_type": {
|
"unsupported_address_type": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{socks5Version, byte(connect), 0, byte(255)},
|
||||||
packet := []byte{socks5Version, byte(connect), 0, byte(255)}
|
|
||||||
return packet
|
|
||||||
},
|
|
||||||
expectedErr: "address type is not supported",
|
expectedErr: "address type is not supported",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -381,59 +366,49 @@ func TestDecodeRequest(t *testing.T) {
|
|||||||
for name, testCase := range testCases {
|
for name, testCase := range testCases {
|
||||||
t.Run(name, func(t *testing.T) {
|
t.Run(name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
packet := testCase.buildPacket()
|
|
||||||
reader := bytes.NewReader(packet)
|
|
||||||
|
|
||||||
req, err := decodeRequest(reader, socks5Version)
|
reader := bytes.NewReader(testCase.packet)
|
||||||
|
|
||||||
|
request, err := decodeRequest(reader, socks5Version)
|
||||||
|
|
||||||
if testCase.expectedErr != "" {
|
if testCase.expectedErr != "" {
|
||||||
assert.ErrorContains(t, err, testCase.expectedErr)
|
assert.ErrorContains(t, err, testCase.expectedErr)
|
||||||
} else {
|
} else {
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
testCase.validate(t, req)
|
testCase.validate(t, request)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestVerifyFirstNegotiation(t *testing.T) {
|
func Test_verifyFirstNegotiation(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
buildPacket func() []byte
|
packet []byte
|
||||||
requiredAuth authMethod
|
requiredAuth authMethod
|
||||||
expectedErr string
|
expectedErr string
|
||||||
}{
|
}{
|
||||||
"version_mismatch": {
|
"version_mismatch": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{4, 2, byte(authNotRequired), byte(authUsernamePassword)},
|
||||||
return []byte{4, 2, byte(authNotRequired), byte(authUsernamePassword)}
|
|
||||||
},
|
|
||||||
requiredAuth: authNotRequired,
|
requiredAuth: authNotRequired,
|
||||||
expectedErr: "version is not supported",
|
expectedErr: "version is not supported",
|
||||||
},
|
},
|
||||||
"no_methods": {
|
"no_methods": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{socks5Version, 0},
|
||||||
return []byte{socks5Version, 0}
|
|
||||||
},
|
|
||||||
requiredAuth: authNotRequired,
|
requiredAuth: authNotRequired,
|
||||||
expectedErr: "no method identifiers",
|
expectedErr: "no method identifiers",
|
||||||
},
|
},
|
||||||
"required_method_not_present": {
|
"required_method_not_present": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{socks5Version, 2, byte(authNotRequired), byte(authGssapi)},
|
||||||
return []byte{socks5Version, 2, byte(authNotRequired), byte(authGssapi)}
|
|
||||||
},
|
|
||||||
requiredAuth: authUsernamePassword,
|
requiredAuth: authUsernamePassword,
|
||||||
expectedErr: "no valid method identifier",
|
expectedErr: "no valid method identifier",
|
||||||
},
|
},
|
||||||
"required_method_present": {
|
"required_method_present": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{socks5Version, 3, byte(authNotRequired), byte(authUsernamePassword), byte(authGssapi)},
|
||||||
return []byte{socks5Version, 3, byte(authNotRequired), byte(authUsernamePassword), byte(authGssapi)}
|
|
||||||
},
|
|
||||||
requiredAuth: authUsernamePassword,
|
requiredAuth: authUsernamePassword,
|
||||||
},
|
},
|
||||||
"no_auth_required": {
|
"no_auth_required": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{socks5Version, 1, byte(authNotRequired)},
|
||||||
return []byte{socks5Version, 1, byte(authNotRequired)}
|
|
||||||
},
|
|
||||||
requiredAuth: authNotRequired,
|
requiredAuth: authNotRequired,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -441,8 +416,8 @@ func TestVerifyFirstNegotiation(t *testing.T) {
|
|||||||
for name, testCase := range testCases {
|
for name, testCase := range testCases {
|
||||||
t.Run(name, func(t *testing.T) {
|
t.Run(name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
packet := testCase.buildPacket()
|
|
||||||
reader := bytes.NewReader(packet)
|
reader := bytes.NewReader(testCase.packet)
|
||||||
|
|
||||||
err := verifyFirstNegotiation(reader, testCase.requiredAuth)
|
err := verifyFirstNegotiation(reader, testCase.requiredAuth)
|
||||||
|
|
||||||
@@ -455,61 +430,54 @@ func TestVerifyFirstNegotiation(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUsernamePasswordSubnegotiate(t *testing.T) {
|
func Test_usernamePasswordSubnegotiate(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
buildPacket func() []byte
|
packet []byte
|
||||||
username string
|
username string
|
||||||
password string
|
password string
|
||||||
expectedErr string
|
expectedErr string
|
||||||
}{
|
}{
|
||||||
"valid_credentials": {
|
"valid_credentials": {
|
||||||
buildPacket: func() []byte {
|
packet: concatBytes(
|
||||||
packet := []byte{authUsernamePasswordSubNegotiation1, 4}
|
[]byte{authUsernamePasswordSubNegotiation1, 4},
|
||||||
packet = append(packet, []byte("user")...)
|
[]byte("user"),
|
||||||
packet = append(packet, 4)
|
[]byte{4},
|
||||||
packet = append(packet, []byte("pass")...)
|
[]byte("pass"),
|
||||||
return packet
|
),
|
||||||
},
|
|
||||||
username: "user",
|
username: "user",
|
||||||
password: "pass",
|
password: "pass",
|
||||||
},
|
},
|
||||||
"version_mismatch": {
|
"version_mismatch": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{2, 4, 'u', 's', 'e', 'r'},
|
||||||
return []byte{2, 4, 'u', 's', 'e', 'r'}
|
|
||||||
},
|
|
||||||
username: "user",
|
username: "user",
|
||||||
password: "pass",
|
password: "pass",
|
||||||
expectedErr: "subnegotiation version not supported",
|
expectedErr: "subnegotiation version not supported",
|
||||||
},
|
},
|
||||||
"wrong_username": {
|
"wrong_username": {
|
||||||
buildPacket: func() []byte {
|
packet: concatBytes(
|
||||||
packet := []byte{authUsernamePasswordSubNegotiation1, 4}
|
[]byte{authUsernamePasswordSubNegotiation1, 4},
|
||||||
packet = append(packet, []byte("fake")...)
|
[]byte("fake"),
|
||||||
packet = append(packet, 4)
|
[]byte{4},
|
||||||
packet = append(packet, []byte("pass")...)
|
[]byte("pass"),
|
||||||
return packet
|
),
|
||||||
},
|
|
||||||
username: "user",
|
username: "user",
|
||||||
password: "pass",
|
password: "pass",
|
||||||
expectedErr: "username not valid",
|
expectedErr: "username received is not valid",
|
||||||
},
|
},
|
||||||
"wrong_password": {
|
"wrong_password": {
|
||||||
buildPacket: func() []byte {
|
packet: concatBytes(
|
||||||
packet := []byte{authUsernamePasswordSubNegotiation1, 4}
|
[]byte{authUsernamePasswordSubNegotiation1, 4},
|
||||||
packet = append(packet, []byte("user")...)
|
[]byte("user"),
|
||||||
packet = append(packet, 4)
|
[]byte{4},
|
||||||
packet = append(packet, []byte("fake")...)
|
[]byte("fake"),
|
||||||
return packet
|
),
|
||||||
},
|
|
||||||
username: "user",
|
username: "user",
|
||||||
password: "pass",
|
password: "pass",
|
||||||
expectedErr: "password not valid",
|
expectedErr: "password not valid",
|
||||||
},
|
},
|
||||||
"truncated_header": {
|
"truncated_header": {
|
||||||
buildPacket: func() []byte {
|
packet: []byte{authUsernamePasswordSubNegotiation1},
|
||||||
return []byte{authUsernamePasswordSubNegotiation1}
|
|
||||||
},
|
|
||||||
username: "user",
|
username: "user",
|
||||||
password: "pass",
|
password: "pass",
|
||||||
expectedErr: "reading header",
|
expectedErr: "reading header",
|
||||||
@@ -519,9 +487,8 @@ func TestUsernamePasswordSubnegotiate(t *testing.T) {
|
|||||||
for name, testCase := range testCases {
|
for name, testCase := range testCases {
|
||||||
t.Run(name, func(t *testing.T) {
|
t.Run(name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
packet := testCase.buildPacket()
|
|
||||||
buffer := &bytes.Buffer{}
|
buffer := bytes.NewBuffer(testCase.packet)
|
||||||
buffer.Write(packet)
|
|
||||||
|
|
||||||
readWriter := struct {
|
readWriter := struct {
|
||||||
io.Reader
|
io.Reader
|
||||||
@@ -542,32 +509,40 @@ func TestUsernamePasswordSubnegotiate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBindDataLength(t *testing.T) {
|
func concatBytes(slices ...[]byte) []byte {
|
||||||
|
var result []byte
|
||||||
|
for _, slice := range slices {
|
||||||
|
result = append(result, slice...)
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_bindDataLength(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
addrType addrType
|
addrType addrType
|
||||||
address string
|
address string
|
||||||
expectedBytes int
|
wantMaxLength uint
|
||||||
}{
|
}{
|
||||||
"ipv4": {
|
"ipv4": {
|
||||||
addrType: ipv4,
|
addrType: ipv4,
|
||||||
address: "127.0.0.1",
|
address: "127.0.0.1",
|
||||||
expectedBytes: 1 + 4 + 2,
|
wantMaxLength: 1 + 4 + 2,
|
||||||
},
|
},
|
||||||
"ipv6": {
|
"ipv6": {
|
||||||
addrType: ipv6,
|
addrType: ipv6,
|
||||||
address: "::1",
|
address: "::1",
|
||||||
expectedBytes: 1 + 16 + 2,
|
wantMaxLength: 1 + 16 + 2,
|
||||||
},
|
},
|
||||||
"domain_short": {
|
"domain_short": {
|
||||||
addrType: domainName,
|
addrType: domainName,
|
||||||
address: "example.com",
|
address: "example.com",
|
||||||
expectedBytes: 1 + 1 + len("example.com") + 2,
|
wantMaxLength: 1 + 1 + uint(len("example.com")) + 2,
|
||||||
},
|
},
|
||||||
"domain_long": {
|
"domain_long": {
|
||||||
addrType: domainName,
|
addrType: domainName,
|
||||||
address: strings.Repeat("a", 100),
|
address: strings.Repeat("a", 100),
|
||||||
expectedBytes: 1 + 1 + 100 + 2,
|
wantMaxLength: 1 + 1 + 100 + 2,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -575,12 +550,12 @@ func TestBindDataLength(t *testing.T) {
|
|||||||
t.Run(name, func(t *testing.T) {
|
t.Run(name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
length := bindDataLength(testCase.addrType, testCase.address)
|
length := bindDataLength(testCase.addrType, testCase.address)
|
||||||
assert.Equal(t, testCase.expectedBytes, length)
|
assert.Equal(t, testCase.wantMaxLength, length)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAuthMethodString(t *testing.T) {
|
func Test_authMethod_String(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
method authMethod
|
method authMethod
|
||||||
@@ -617,7 +592,7 @@ func TestAuthMethodString(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCmdTypeString(t *testing.T) {
|
func Test_cmdType_String(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
testCases := map[string]struct {
|
testCases := map[string]struct {
|
||||||
cmd cmdType
|
cmd cmdType
|
||||||
@@ -649,34 +624,3 @@ func TestCmdTypeString(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestParseAddress(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
testCases := map[string]struct {
|
|
||||||
address string
|
|
||||||
expectedIP string
|
|
||||||
expectedErr string
|
|
||||||
}{
|
|
||||||
"ipv4": {
|
|
||||||
address: "127.0.0.1",
|
|
||||||
expectedIP: "127.0.0.1",
|
|
||||||
},
|
|
||||||
"ipv6": {
|
|
||||||
address: "::1",
|
|
||||||
expectedIP: "::1",
|
|
||||||
},
|
|
||||||
"domain": {
|
|
||||||
address: "example.com",
|
|
||||||
expectedErr: "parsing IP address",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for name, testCase := range testCases {
|
|
||||||
t.Run(name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
if testCase.expectedErr == "" {
|
|
||||||
assert.True(t, strings.Contains(testCase.address, testCase.expectedIP) || testCase.address == testCase.expectedIP)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ func usernamePasswordSubnegotiate(conn io.ReadWriter, username, password string)
|
|||||||
return fmt.Errorf("reading username bytes: %w", err)
|
return fmt.Errorf("reading username bytes: %w", err)
|
||||||
} else if username != string(usernameBytes) {
|
} else if username != string(usernameBytes) {
|
||||||
_, _ = conn.Write([]byte{version, status})
|
_, _ = conn.Write([]byte{version, status})
|
||||||
return fmt.Errorf("username not valid: %s", string(usernameBytes))
|
return fmt.Errorf("username received is not valid")
|
||||||
}
|
}
|
||||||
|
|
||||||
const passwordHeaderLength = 1
|
const passwordHeaderLength = 1
|
||||||
|
|||||||
Reference in New Issue
Block a user