review tests and fix AI found issues

This commit is contained in:
Quentin McGaw
2026-05-19 13:36:02 +00:00
parent 4a0e8afd3b
commit 5cd7bc7f74
5 changed files with 111 additions and 161 deletions
+2 -2
View File
@@ -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
View File
@@ -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()
} }
+2 -1
View File
@@ -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
View File
@@ -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)
}
})
}
}
+1 -1
View File
@@ -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