From 5cd7bc7f7488437893e9cf4457eb343935766415 Mon Sep 17 00:00:00 2001 From: Quentin McGaw Date: Tue, 19 May 2026 13:36:02 +0000 Subject: [PATCH] review tests and fix AI found issues --- internal/socks5/response.go | 4 +- internal/socks5/server.go | 29 ++-- internal/socks5/socks5.go | 3 +- internal/socks5/socks5_test.go | 234 +++++++++++----------------- internal/socks5/usernamepassword.go | 2 +- 5 files changed, 111 insertions(+), 161 deletions(-) diff --git a/internal/socks5/response.go b/internal/socks5/response.go index decd0bae..4df7ebee 100644 --- a/internal/socks5/response.go +++ b/internal/socks5/response.go @@ -91,14 +91,14 @@ func encodeBindData(addrType addrType, address string, port uint16) ( return data, nil } -func bindDataLength(addrType addrType, address string) (maxLength int) { +func bindDataLength(addrType addrType, address string) (maxLength uint) { maxLength++ // address type switch addrType { case ipv4: maxLength += net.IPv4len case domainName: maxLength++ // domain name length - maxLength += len([]byte(address)) + maxLength += uint(len([]byte(address))) case ipv6: maxLength += net.IPv6len default: diff --git a/internal/socks5/server.go b/internal/socks5/server.go index a6851400..26f6f250 100644 --- a/internal/socks5/server.go +++ b/internal/socks5/server.go @@ -8,7 +8,7 @@ import ( "sync/atomic" ) -type Server struct { +type server struct { username string password string address string @@ -23,8 +23,8 @@ type Server struct { stopping atomic.Bool } -func newServer(settings Settings) *Server { - return &Server{ +func newServer(settings Settings) *server { + return &server{ username: settings.Username, password: settings.Password, 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()) config := &net.ListenConfig{} 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 } -func (s *Server) runServer(ready chan<- struct{}, +func (s *server) runServer(ready chan<- struct{}, runErrCh chan<- error, done chan<- struct{}, ) { close(ready) @@ -69,7 +69,7 @@ func (s *Server) runServer(ready chan<- struct{}, connection, err := s.listener.Accept() if err != nil { if !s.stopping.Load() { - _ = s.Stop() + _ = s.stop() runErrCh <- fmt.Errorf("accepting connection: %w", err) } 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.listening.Store(false) - err = s.listener.Close() - s.socksConnCancel() // stop ongoing socks connections - <-s.done // wait for run goroutine to finish + err = s.stop() + <-s.done // wait for run goroutine to finish s.stopping.Store(false) 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() { return s.listener.Addr() } diff --git a/internal/socks5/socks5.go b/internal/socks5/socks5.go index 0db95cfb..a92a778a 100644 --- a/internal/socks5/socks5.go +++ b/internal/socks5/socks5.go @@ -132,7 +132,8 @@ func (c *socksConn) handleRequest(ctx context.Context) error { 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() { _, err := io.Copy(c.clientConn, destinationConn) if err != nil { diff --git a/internal/socks5/socks5_test.go b/internal/socks5/socks5_test.go index 4545e7d9..2174e015 100644 --- a/internal/socks5/socks5_test.go +++ b/internal/socks5/socks5_test.go @@ -2,7 +2,6 @@ package socks5 import ( "bytes" - "context" "encoding/binary" "io" "net" @@ -17,8 +16,8 @@ import ( type noopLogger struct{} -func (noopLogger) Infof(string, ...interface{}) {} -func (noopLogger) Warnf(string, ...interface{}) {} +func (noopLogger) Infof(string, ...any) {} +func (noopLogger) Warnf(string, ...any) {} func TestServerProxy(t *testing.T) { t.Parallel() @@ -38,10 +37,10 @@ func TestServerProxy(t *testing.T) { t.Parallel() // 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) - backendConnCh := make(chan net.Conn, 1) + backendConnCh := make(chan net.Conn) go func() { conn, err := backendListener.Accept() if err != nil { @@ -56,7 +55,7 @@ func TestServerProxy(t *testing.T) { Address: "127.0.0.1:0", Logger: noopLogger{}, }) - _, err = server.Start(context.Background()) + _, err = server.Start(t.Context()) require.NoError(t, err) t.Cleanup(func() { _ = server.Stop() @@ -107,7 +106,7 @@ func dialSOCKS5(t *testing.T, proxyAddr, targetAddr, username, password string) targetPort, err := strconv.Atoi(portStr) 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) var method authMethod @@ -180,11 +179,11 @@ func dialSOCKS5(t *testing.T, proxyAddr, targetAddr, username, password string) return conn } -func TestNew(t *testing.T) { +func Test_newServer(t *testing.T) { t.Parallel() testCases := map[string]struct { settings Settings - expected *Server + expected *server }{ "with_auth": { settings: Settings{ @@ -193,7 +192,7 @@ func TestNew(t *testing.T) { Address: "127.0.0.1:1080", Logger: nil, }, - expected: &Server{ + expected: &server{ username: "user", password: "pass", address: "127.0.0.1:1080", @@ -205,7 +204,7 @@ func TestNew(t *testing.T) { Address: "127.0.0.1:1080", Logger: nil, }, - expected: &Server{ + expected: &server{ address: "127.0.0.1:1080", 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() ctrl := gomock.NewController(t) logger := NewMockLogger(ctrl) @@ -236,7 +235,7 @@ func TestStartStop(t *testing.T) { Logger: logger, }) - runErr, startErr := server.Start(context.Background()) + runErr, startErr := server.Start(t.Context()) require.NoError(t, startErr) select { @@ -252,7 +251,7 @@ func TestStartStop(t *testing.T) { require.NoError(t, err) } -func TestEncodeBindData(t *testing.T) { +func Test_encodeBindData(t *testing.T) { t.Parallel() testCases := map[string]struct { addrType addrType @@ -319,61 +318,47 @@ func TestEncodeBindData(t *testing.T) { } } -func TestDecodeRequest(t *testing.T) { +func Test_decodeRequest(t *testing.T) { t.Parallel() testCases := map[string]struct { - buildPacket func() []byte + packet []byte expectedErr string validate func(*testing.T, request) }{ "ipv4_valid": { - buildPacket: func() []byte { - packet := []byte{socks5Version, byte(connect), 0, byte(ipv4)} - packet = append(packet, 127, 0, 0, 1) - packet = append(packet, byte(0x1f), byte(0x90)) - return packet - }, - validate: func(t *testing.T, req request) { + packet: []byte{socks5Version, byte(connect), 0, byte(ipv4), 127, 0, 0, 1, byte(0x1f), byte(0x90)}, + validate: func(t *testing.T, request request) { t.Helper() - assert.Equal(t, connect, req.command) - assert.Equal(t, "127.0.0.1", req.destination) - assert.Equal(t, uint16(8080), req.port) - assert.Equal(t, ipv4, req.addressType) + assert.Equal(t, connect, request.command) + assert.Equal(t, "127.0.0.1", request.destination) + assert.Equal(t, uint16(8080), request.port) + assert.Equal(t, ipv4, request.addressType) }, }, "domain_name_valid": { - buildPacket: func() []byte { - packet := []byte{socks5Version, byte(connect), 0, byte(domainName)} - domain := "example.com" - packet = append(packet, byte(len(domain))) - packet = append(packet, []byte(domain)...) - packet = append(packet, byte(0x00), byte(0x50)) - return packet - }, - validate: func(t *testing.T, req request) { + packet: concatBytes( + []byte{socks5Version, byte(connect), 0, byte(domainName)}, + []byte{byte(len("example.com"))}, + []byte("example.com"), + []byte{0x00, 0x50}, + ), + validate: func(t *testing.T, request request) { t.Helper() - assert.Equal(t, "example.com", req.destination) - assert.Equal(t, uint16(80), req.port) - assert.Equal(t, domainName, req.addressType) + assert.Equal(t, "example.com", request.destination) + assert.Equal(t, uint16(80), request.port) + assert.Equal(t, domainName, request.addressType) }, }, "version_mismatch": { - buildPacket: func() []byte { - return []byte{4, byte(connect), 0, byte(ipv4), 127, 0, 0, 1, 0, 0} - }, + packet: []byte{4, byte(connect), 0, byte(ipv4), 127, 0, 0, 1, 0, 0}, expectedErr: "version is not supported", }, "truncated_header": { - buildPacket: func() []byte { - return []byte{socks5Version, byte(connect)} - }, + packet: []byte{socks5Version, byte(connect)}, expectedErr: "reading header", }, "unsupported_address_type": { - buildPacket: func() []byte { - packet := []byte{socks5Version, byte(connect), 0, byte(255)} - return packet - }, + packet: []byte{socks5Version, byte(connect), 0, byte(255)}, expectedErr: "address type is not supported", }, } @@ -381,59 +366,49 @@ func TestDecodeRequest(t *testing.T) { for name, testCase := range testCases { t.Run(name, func(t *testing.T) { 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 != "" { assert.ErrorContains(t, err, testCase.expectedErr) } else { 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() testCases := map[string]struct { - buildPacket func() []byte + packet []byte requiredAuth authMethod expectedErr string }{ "version_mismatch": { - buildPacket: func() []byte { - return []byte{4, 2, byte(authNotRequired), byte(authUsernamePassword)} - }, + packet: []byte{4, 2, byte(authNotRequired), byte(authUsernamePassword)}, requiredAuth: authNotRequired, expectedErr: "version is not supported", }, "no_methods": { - buildPacket: func() []byte { - return []byte{socks5Version, 0} - }, + packet: []byte{socks5Version, 0}, requiredAuth: authNotRequired, expectedErr: "no method identifiers", }, "required_method_not_present": { - buildPacket: func() []byte { - return []byte{socks5Version, 2, byte(authNotRequired), byte(authGssapi)} - }, + packet: []byte{socks5Version, 2, byte(authNotRequired), byte(authGssapi)}, requiredAuth: authUsernamePassword, expectedErr: "no valid method identifier", }, "required_method_present": { - buildPacket: func() []byte { - return []byte{socks5Version, 3, byte(authNotRequired), byte(authUsernamePassword), byte(authGssapi)} - }, + packet: []byte{socks5Version, 3, byte(authNotRequired), byte(authUsernamePassword), byte(authGssapi)}, requiredAuth: authUsernamePassword, }, "no_auth_required": { - buildPacket: func() []byte { - return []byte{socks5Version, 1, byte(authNotRequired)} - }, + packet: []byte{socks5Version, 1, byte(authNotRequired)}, requiredAuth: authNotRequired, }, } @@ -441,8 +416,8 @@ func TestVerifyFirstNegotiation(t *testing.T) { for name, testCase := range testCases { t.Run(name, func(t *testing.T) { t.Parallel() - packet := testCase.buildPacket() - reader := bytes.NewReader(packet) + + reader := bytes.NewReader(testCase.packet) 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() testCases := map[string]struct { - buildPacket func() []byte + packet []byte username string password string expectedErr string }{ "valid_credentials": { - buildPacket: func() []byte { - packet := []byte{authUsernamePasswordSubNegotiation1, 4} - packet = append(packet, []byte("user")...) - packet = append(packet, 4) - packet = append(packet, []byte("pass")...) - return packet - }, + packet: concatBytes( + []byte{authUsernamePasswordSubNegotiation1, 4}, + []byte("user"), + []byte{4}, + []byte("pass"), + ), username: "user", password: "pass", }, "version_mismatch": { - buildPacket: func() []byte { - return []byte{2, 4, 'u', 's', 'e', 'r'} - }, + packet: []byte{2, 4, 'u', 's', 'e', 'r'}, username: "user", password: "pass", expectedErr: "subnegotiation version not supported", }, "wrong_username": { - buildPacket: func() []byte { - packet := []byte{authUsernamePasswordSubNegotiation1, 4} - packet = append(packet, []byte("fake")...) - packet = append(packet, 4) - packet = append(packet, []byte("pass")...) - return packet - }, + packet: concatBytes( + []byte{authUsernamePasswordSubNegotiation1, 4}, + []byte("fake"), + []byte{4}, + []byte("pass"), + ), username: "user", password: "pass", - expectedErr: "username not valid", + expectedErr: "username received is not valid", }, "wrong_password": { - buildPacket: func() []byte { - packet := []byte{authUsernamePasswordSubNegotiation1, 4} - packet = append(packet, []byte("user")...) - packet = append(packet, 4) - packet = append(packet, []byte("fake")...) - return packet - }, + packet: concatBytes( + []byte{authUsernamePasswordSubNegotiation1, 4}, + []byte("user"), + []byte{4}, + []byte("fake"), + ), username: "user", password: "pass", expectedErr: "password not valid", }, "truncated_header": { - buildPacket: func() []byte { - return []byte{authUsernamePasswordSubNegotiation1} - }, + packet: []byte{authUsernamePasswordSubNegotiation1}, username: "user", password: "pass", expectedErr: "reading header", @@ -519,9 +487,8 @@ func TestUsernamePasswordSubnegotiate(t *testing.T) { for name, testCase := range testCases { t.Run(name, func(t *testing.T) { t.Parallel() - packet := testCase.buildPacket() - buffer := &bytes.Buffer{} - buffer.Write(packet) + + buffer := bytes.NewBuffer(testCase.packet) readWriter := struct { 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() testCases := map[string]struct { addrType addrType address string - expectedBytes int + wantMaxLength uint }{ "ipv4": { addrType: ipv4, address: "127.0.0.1", - expectedBytes: 1 + 4 + 2, + wantMaxLength: 1 + 4 + 2, }, "ipv6": { addrType: ipv6, address: "::1", - expectedBytes: 1 + 16 + 2, + wantMaxLength: 1 + 16 + 2, }, "domain_short": { addrType: domainName, address: "example.com", - expectedBytes: 1 + 1 + len("example.com") + 2, + wantMaxLength: 1 + 1 + uint(len("example.com")) + 2, }, "domain_long": { addrType: domainName, 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.Parallel() 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() testCases := map[string]struct { 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() testCases := map[string]struct { 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) - } - }) - } -} diff --git a/internal/socks5/usernamepassword.go b/internal/socks5/usernamepassword.go index a0671013..34be9835 100644 --- a/internal/socks5/usernamepassword.go +++ b/internal/socks5/usernamepassword.go @@ -31,7 +31,7 @@ func usernamePasswordSubnegotiate(conn io.ReadWriter, username, password string) return fmt.Errorf("reading username bytes: %w", err) } else if username != string(usernameBytes) { _, _ = 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